Rejection sampling and residual distributions
Overview
Section titled “Overview”| Module | M07.6 · build · Python · Pass 6 · 2 to 3 h |
| You build | python/tinyllm/prob/rejection.py: rejection_accept, acceptance_probability, residual_distribution, speculative_step, rejection_sample |
| Contract | course/contracts/py/tinyllm/prob/rejection.pyi · the draw order: spec/sampling.md (draw accounting) |
| Tests | course/tests/M07.6/ (what they check: section 4) |
| Needs | M07.1 (sample_categorical, the inverse CDF, and UniformSource) · reading: M07.0 random variables, M06.3 PCG32 (your uniforms) |
| Used by | L8.6 speculative decoding verifies every draft token with speculative_step, and L10.8 ports it to Rust |
| Milestone | MS-P6 (Pass 6 gate: every math module of the pass checks green) |
| Optional depth | Devroye, Non-Uniform Random Variate Generation (1986), ch. 2.3; Leviathan, Kalman, and Matias, “Fast Inference from Transformers via Speculative Decoding” (2023), appendix A.1; Chen et al., “Accelerating Large Language Model Decoding with Speculative Sampling” (2023) |
Key Takeaways
Section titled “Key Takeaways”- Keeping a draft with probability keeps exactly of every token’s mass; the total kept is (
test_acceptance_rate_matches_acceptance_probability). - What rejection removes is exactly the residual , so resampling from its normalized form restores token for token: the output of
speculative_stepis distributed as for every (test_speculative_output_is_exactly_p). - “Accept when ” with uniform on has probability exactly ; "" counts one extra point and accepts tokens never produces (
test_accept_is_strict). - Von Neumann’s sampler accepts a proposal with probability ; each try succeeds with probability , so the cost is a geometric number of tries with mean (
test_rejection_sample_tries_are_geometric).
How to work this chapter
Section titled “How to work this chapter”ol start M07.6 # stubs rejection.py into your repool tests M07.6 # read the test catalog first: rung R0, you write no tests hereol check M07.6 # exit code is the verdictol check M07.6 --ref-deps # only if you skipped M07.1ol diff M07.6 # after passing: your code against the reference1. Why now
Section titled “1. Why now”Your sampler (L8.1) emits one token per forward pass of the model, and each forward pass reads every weight once. In Part 8 you will make decoding faster without changing what it says: a cheap draft (an n-gram table, a prompt lookup, a small model) proposes several tokens, and the target model checks all of them in one forward pass (L8.6). Greedy decoding only needs “did the draft guess the argmax”. Sampled decoding needs more: the accepted text must be distributed exactly as if the target had sampled it alone, or the speedup silently changes the model’s output distribution, which no eval would attribute to the cache. This module is the piece of probability that makes that exact: an acceptance test and a residual distribution, proved by enumeration before L8.6 builds on them.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| number of ids (the vocabulary) | integer | |
| the target distribution (the big model’s sampler output) | float64[n], sums to 1 | |
| the proposal distribution (the draft) | float64[n], sums to 1 | |
| a draft token, drawn from | integer in | |
| independent uniforms on | float | |
| the acceptance ratio of | float | |
| the acceptance probability | float in | |
| total variation distance | float in | |
| the residual distribution, | float64[n] | |
| an envelope: for every | float |
2.1 Von Neumann’s rejection sampler
Section titled “2.1 Von Neumann’s rejection sampler”You can sample cheaply but want samples of . If for every , repeat: draw , draw , and accept when . On one try,
Conditioned on acceptance, has probability : exactly the target. Each try is an independent coin with success probability , so the number of tries is geometric, with mean and variance . A loose envelope costs time, never correctness; an envelope that is too small ( for some ) caps the acceptance of at 1 and under-samples it, which is why the contract checks it.
2.2 The acceptance test, exactly
Section titled “2.2 The acceptance test, exactly”For uniform on and any , . The comparison must be strict: is the same for continuous , but floating-point uniforms are discrete (M06.3 gives multiples of ), and the point with would accept a token the target never produces. With , every is below it: an under-proposed token is always kept.
2.3 Accept, or draw from the residual
Section titled “2.3 Accept, or draw from the residual”Speculative decoding has no retry loop: the target’s forward pass already happened, so a rejected position must produce a token right away. Keep with probability . The kept mass of token is
and the total kept mass is . Since and both distributions sum to 1, : the closer the draft, the more it is accepted. What is missing from after acceptance is , which sums to . Normalizing gives the residual.
2.4 The theorem
Section titled “2.4 The theorem”Draw , accept with probability , otherwise draw . Then
Nothing was assumed about beyond being a distribution: a bad draft makes the method slow, never wrong. When , and the residual is undefined, but then and it is never sampled; the contract returns so the function stays total.
2.5 Determinism
Section titled “2.5 Determinism”The sums in and are left-to-right loops over ascending ids in float64, and , come from the request’s generator in the order L8.6 fixes. The Rust engine (L10.8) repeats the same arithmetic, so both accept and reject the same drafts for the same seed. numpy’s np.sum sums pairwise and can differ in the last bit, which moves a boundary of the residual’s CDF.
3. Worked example by hand
Section titled “3. Worked example by hand”Four ids, (the target), (the draft).
| id | kept mass | |||||
|---|---|---|---|---|---|---|
| 0 | 1/2 | 1/4 | 2 | 1 | 1/4 | 1/4 |
| 1 | 1/4 | 1/2 | 1/2 | 1/2 | 1/4 | 0 |
| 2 | 1/4 | 0 | (never proposed) | 0 | 1/4 | |
| 3 | 0 | 1/4 | 0 | 0 | 0 | 0 |
The acceptance probability is ; the residual mass is , and the residual is . The output distribution:
- id 0: kept , plus rejected mass times residual : .
- id 1: kept , residual 0: .
- id 2: never proposed, residual: .
- id 3: proposed a quarter of the time and always rejected: 0.
That is . One concrete draw: the draft proposes , rejects it, and walks the residual’s CDF to id 2. With instead, id 1 is kept and is not looked at. These are the first cases in section 4: test_hand_example, test_hand_example_output_is_p, and test_speculative_accepts_without_resampling.
4. The interface
Section titled “4. The interface”def rejection_accept(p_x: float, q_x: float, u: float) -> bool # u < p_x / q_xdef acceptance_probability(p: ArrayLike, q: ArrayLike) -> float # sum min(p, q)def residual_distribution(p: ArrayLike, q: ArrayLike) -> NDArray # normalize(max(0, p - q)); p if p == qdef speculative_step(p, q, x: int, u_accept: float, u_resample: float) -> tuple[int, bool]def rejection_sample(p, q, m: float, rng: UniformSource, max_tries: int = 10000) -> tuple[int, int]speculative_step resamples with your M07.1 sample_categorical. rejection_sample draws exactly two uniforms per try, proposal first. The contract has every error rule.
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example | unit | section 3’s ratios, , residual | you and the test agree on the definitions |
test_hand_example_output_is_p | statistical | section 3’s output distribution, enumerated exactly | the theorem on numbers you can check by hand |
test_accept_is_strict | boundary | rejects, never accepts | discrete uniforms make wrong |
test_accept_always_when_target_dominates | boundary | accepts every | under-proposed tokens are always kept |
test_accept_rejects_bad_arguments | boundary | , , raise | a silent answer biases the output |
test_acceptance_rate_matches_acceptance_probability | statistical | enumerated acceptance | L8.6 reports it as its acceptance rate |
test_residual_is_a_distribution | property | non-negative, sums to 1, | sample_categorical demands a distribution |
test_residual_when_p_equals_q | boundary | gives a copy of | no 0/0, no aliasing of the caller’s array |
test_residual_disjoint_supports_is_p | boundary | gives residual | a useless draft is slow, not wrong |
test_residual_sum_is_ascending | unit | the normalizer is a left-to-right loop | bit parity with the Rust port (L10.8) |
test_shapes_and_distributions_are_checked | boundary | mismatched shapes and non-distributions raise | vocabularies of draft and target must match |
test_speculative_output_is_exactly_p | statistical | output distribution exactly, random in eighths | the theorem L8.6 relies on |
test_speculative_chi_square | statistical | 20000 seeded draws match (chi-square) | the same law on arbitrary distributions |
test_speculative_accepts_without_resampling | unit | an accepted draft ignores | the draw order is part of the parity contract |
test_speculative_rejects_bad_drafts | boundary | outside the vocabulary or raise | draft and proposal disagreeing is a bug |
test_rejection_sample_first_try_enumerated | statistical | one try accepts with probability , exactly | the von Neumann law |
test_rejection_sample_tries_are_geometric | statistical | mean tries within 4 standard errors | the cost of a loose envelope |
test_rejection_sample_draw_order | unit | proposal uniform first, then acceptance; two per try | reproducible across languages |
test_rejection_sample_checks_the_envelope | boundary | , , and no acceptance in max_tries raise | a wrong envelope silently biases |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. accepting on | a token with is accepted when | test_accept_is_strict (mutant s01) |
| 2. the ratio upside down, or “capped” by inverting it | over-proposed tokens kept, under-proposed ones dropped | test_accept_always_when_target_dominates (mutants s02, s12) |
| 3. residual or | resamples the tokens the draft over-proposed | test_residual_is_a_distribution (mutants s03, s05) |
| 4. forgetting to normalize the residual | sample_categorical rejects it, or the CDF ends below 1 | test_residual_is_a_distribution (mutant s04) |
| 5. on rejection, resampling from , from , or returning the draft | output no longer : accepted tokens counted twice | test_speculative_output_is_exactly_p (mutants s06, s07, s13) |
| 6. handled as 0/0, or returning the caller’s array | NaN, or a later in-place edit corrupts the target | test_residual_when_p_equals_q (mutants s08, s14) |
| 7. von Neumann without , or without checking it | the accepted tokens follow -ish weights, not | test_rejection_sample_first_try_enumerated (mutant s15), test_rejection_sample_checks_the_envelope (mutant s19) |
8. summing with numpy’s pairwise np.sum | last-bit differences against the Rust port | test_residual_sum_is_ascending (mutant s09) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | M07.1 | sample_categorical draws the residual; UniformSource is the generator type |
| Back | M07.0 | uniform random variables and independence, behind every probability above |
| Forward | L8.6 | speculative decoding: one speculative_step per draft position, acceptance_probability as the reported acceptance rate |
| Forward | L10.8 | the Rust engine’s speculative decoding, held to the same accept and reject decisions |
If you skip this module, L8.6 cannot check: ol check L8.6 stops with L8.6 needs M07.6, and --ref-deps substitutes the reference.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
speculative_step | vLLM RejectionSampler | batched acceptance over a whole draft tree on the GPU, recovered tokens from the residual | vllm/v1/sample/rejection_sampler.py |
residual_distribution | Hugging Face assisted generation | the same residual for model drafts (_speculative_sampling) | transformers/generation/utils.py |
rejection_sample | NumPy’s Gamma and binomial samplers | rejection with tight squeeze functions to avoid evaluating the target | numpy/random/src/distributions/distributions.c |