Skip to content

Rejection sampling and residual distributions

ModuleM07.6 · build · Python · Pass 6 · 2 to 3 h
You buildpython/tinyllm/prob/rejection.py: rejection_accept, acceptance_probability, residual_distribution, speculative_step, rejection_sample
Contractcourse/contracts/py/tinyllm/prob/rejection.pyi · the draw order: spec/sampling.md (draw accounting)
Testscourse/tests/M07.6/ (what they check: section 4)
NeedsM07.1 (sample_categorical, the inverse CDF, and UniformSource) · reading: M07.0 random variables, M06.3 PCG32 (your uniforms)
Used byL8.6 speculative decoding verifies every draft token with speculative_step, and L10.8 ports it to Rust
MilestoneMS-P6 (Pass 6 gate: every math module of the pass checks green)
Optional depthDevroye, 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)
  • Keeping a draft x∼qx \sim q with probability min⁡(1,px/qx)\min(1, p_x/q_x) keeps exactly min⁡(px,qx)\min(p_x, q_x) of every token’s mass; the total kept is ∑xmin⁡(px,qx)=1−TV(p,q)\sum_x \min(p_x, q_x) = 1 - \mathrm{TV}(p, q) (test_acceptance_rate_matches_acceptance_probability).
  • What rejection removes is exactly the residual max⁡(0,p−q)\max(0, p - q), so resampling from its normalized form restores pp token for token: the output of speculative_step is distributed as pp for every qq (test_speculative_output_is_exactly_p).
  • “Accept when u<ru < r” with uu uniform on [0,1)[0, 1) has probability exactly min⁡(1,r)\min(1, r); "u≤ru \le r" counts one extra point and accepts tokens pp never produces (test_accept_is_strict).
  • Von Neumann’s sampler accepts a proposal with probability px/(mqx)p_x / (m q_x); each try succeeds with probability 1/m1/m, so the cost is a geometric number of tries with mean mm (test_rejection_sample_tries_are_geometric).
Terminal window
ol start M07.6 # stubs rejection.py into your repo
ol tests M07.6 # read the test catalog first: rung R0, you write no tests here
ol check M07.6 # exit code is the verdict
ol check M07.6 --ref-deps # only if you skipped M07.1
ol diff M07.6 # after passing: your code against the reference

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.

SymbolMeaningType / shape
nnnumber of ids (the vocabulary)integer
p∈Δn−1p \in \Delta^{n-1}the target distribution (the big model’s sampler output)float64[n], sums to 1
q∈Δn−1q \in \Delta^{n-1}the proposal distribution (the draft)float64[n], sums to 1
xxa draft token, drawn from qqinteger in [0,n)[0, n)
u,u′u, u'independent uniforms on [0,1)[0, 1)float
rx=px/qxr_x = p_x / q_xthe acceptance ratio of xxfloat ≥0\ge 0
α=∑imin⁡(pi,qi)\alpha = \sum_i \min(p_i, q_i)the acceptance probabilityfloat in [0,1][0, 1]
TV(p,q)=12∑i∣pi−qi∣\mathrm{TV}(p, q) = \tfrac12 \sum_i \lvert p_i - q_i\rverttotal variation distancefloat in [0,1][0, 1]
res(p,q)i=max⁡(0,pi−qi)/Z\mathrm{res}(p, q)_i = \max(0, p_i - q_i) / Zthe residual distribution, Z=∑imax⁡(0,pi−qi)Z = \sum_i \max(0, p_i - q_i)float64[n]
mman envelope: pi≤mqip_i \le m q_i for every iifloat ≥1\ge 1

You can sample qq cheaply but want samples of pp. If pi≤m qip_i \le m\, q_i for every ii, repeat: draw x∼qx \sim q, draw uu, and accept xx when u<px/(m qx)u < p_x / (m\, q_x). On one try,

P(propose x and accept)=qx⋅pxm qx=pxm,P(accept)=∑xpxm=1m.P(\text{propose } x \text{ and accept}) = q_x \cdot \frac{p_x}{m\, q_x} = \frac{p_x}{m}, \qquad P(\text{accept}) = \sum_x \frac{p_x}{m} = \frac{1}{m}.

Conditioned on acceptance, xx has probability (px/m)/(1/m)=px(p_x/m)/(1/m) = p_x: exactly the target. Each try is an independent coin with success probability 1/m1/m, so the number of tries is geometric, with mean mm and variance m(m−1)m(m - 1). A loose envelope costs time, never correctness; an envelope that is too small (px>mqxp_x > m q_x for some xx) caps the acceptance of xx at 1 and under-samples it, which is why the contract checks it.

For uu uniform on [0,1)[0, 1) and any r≥0r \ge 0, P(u<r)=min⁡(1,r)P(u < r) = \min(1, r). The comparison must be strict: P(u≤r)P(u \le r) is the same for continuous uu, but floating-point uniforms are discrete (M06.3 gives multiples of 2−532^{-53}), and the point u=0u = 0 with px=0p_x = 0 would accept a token the target never produces. With r=px/qx≥1r = p_x / q_x \ge 1, every u<1u < 1 is below it: an under-proposed token is always kept.

Speculative decoding has no retry loop: the target’s forward pass already happened, so a rejected position must produce a token right away. Keep x∼qx \sim q with probability min⁡(1,px/qx)\min(1, p_x/q_x). The kept mass of token tt is

qtmin⁡ ⁣(1,ptqt)=min⁡(pt,qt),q_t \min\!\left(1, \frac{p_t}{q_t}\right) = \min(p_t, q_t),

and the total kept mass is α=∑tmin⁡(pt,qt)\alpha = \sum_t \min(p_t, q_t). Since min⁡(a,b)=12(a+b−∣a−b∣)\min(a, b) = \tfrac12(a + b - \lvert a - b\rvert) and both distributions sum to 1, α=1−TV(p,q)\alpha = 1 - \mathrm{TV}(p, q): the closer the draft, the more it is accepted. What is missing from pp after acceptance is pt−min⁡(pt,qt)=max⁡(0,pt−qt)p_t - \min(p_t, q_t) = \max(0, p_t - q_t), which sums to Z=1−αZ = 1 - \alpha. Normalizing gives the residual.

Draw x∼qx \sim q, accept with probability min⁡(1,px/qx)\min(1, p_x/q_x), otherwise draw t∼res(p,q)t \sim \mathrm{res}(p, q). Then

P(output=t)=min⁡(pt,qt)⏟accepted t+(1−α)⏟rejected⋅max⁡(0,pt−qt)1−α=min⁡(pt,qt)+max⁡(0,pt−qt)=pt.P(\text{output} = t) = \underbrace{\min(p_t, q_t)}_{\text{accepted } t} + \underbrace{(1 - \alpha)}_{\text{rejected}} \cdot \frac{\max(0, p_t - q_t)}{1 - \alpha} = \min(p_t, q_t) + \max(0, p_t - q_t) = p_t.

Nothing was assumed about qq beyond being a distribution: a bad draft makes the method slow, never wrong. When p=qp = q, Z=0Z = 0 and the residual is undefined, but then α=1\alpha = 1 and it is never sampled; the contract returns pp so the function stays total.

The sums in α\alpha and ZZ are left-to-right loops over ascending ids in float64, and uu, u′u' 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.

Four ids, p=[12,14,14,0]p = [\tfrac12, \tfrac14, \tfrac14, 0] (the target), q=[14,12,0,14]q = [\tfrac14, \tfrac12, 0, \tfrac14] (the draft).

id ttptp_tqtq_trt=pt/qtr_t = p_t/q_tP(keep∣x=t)P(\text{keep} \mid x = t)kept mass min⁡(pt,qt)\min(p_t, q_t)max⁡(0,pt−qt)\max(0, p_t - q_t)
01/21/4211/41/4
11/41/21/21/21/40
21/40(never proposed)01/4
301/40000

The acceptance probability is α=14+14=12\alpha = \tfrac14 + \tfrac14 = \tfrac12; the residual mass is Z=12Z = \tfrac12, and the residual is [12,0,12,0][\tfrac12, 0, \tfrac12, 0]. The output distribution:

  • id 0: kept 14\tfrac14, plus rejected mass 12\tfrac12 times residual 12\tfrac12: 12\tfrac12.
  • id 1: kept 14\tfrac14, residual 0: 14\tfrac14.
  • id 2: never proposed, residual: 12⋅12=14\tfrac12 \cdot \tfrac12 = \tfrac14.
  • id 3: proposed a quarter of the time and always rejected: 0.

That is pp. One concrete draw: the draft proposes x=1x = 1, u=0.75≥r1=0.5u = 0.75 \ge r_1 = 0.5 rejects it, and u′=0.75u' = 0.75 walks the residual’s CDF [0.5,0.5,1.0,1.0][0.5, 0.5, 1.0, 1.0] to id 2. With u=0.25u = 0.25 instead, id 1 is kept and u′u' 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.

python/tinyllm/prob/rejection.py
def rejection_accept(p_x: float, q_x: float, u: float) -> bool # u < p_x / q_x
def 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 == q
def 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.

TestKINDChecksWhy it matters downstream
test_hand_exampleunitsection 3’s ratios, α=1/2\alpha = 1/2, residual [1/2,0,1/2,0][1/2, 0, 1/2, 0]you and the test agree on the definitions
test_hand_example_output_is_pstatisticalsection 3’s output distribution, enumerated exactlythe theorem on numbers you can check by hand
test_accept_is_strictboundaryu=ru = r rejects, px=0p_x = 0 never acceptsdiscrete uniforms make ≤\le wrong
test_accept_always_when_target_dominatesboundaryr≥1r \ge 1 accepts every u<1u < 1under-proposed tokens are always kept
test_accept_rejects_bad_argumentsboundaryqx=0q_x = 0, u∉[0,1)u \notin [0, 1), px<0p_x < 0 raisea silent answer biases the output
test_acceptance_rate_matches_acceptance_probabilitystatisticalenumerated acceptance =∑min⁡(p,q)= \sum \min(p, q)L8.6 reports it as its acceptance rate
test_residual_is_a_distributionpropertynon-negative, sums to 1, Z=1−αZ = 1 - \alphasample_categorical demands a distribution
test_residual_when_p_equals_qboundaryp=qp = q gives a copy of ppno 0/0, no aliasing of the caller’s array
test_residual_disjoint_supports_is_pboundaryα=0\alpha = 0 gives residual ppa useless draft is slow, not wrong
test_residual_sum_is_ascendingunitthe normalizer is a left-to-right loopbit parity with the Rust port (L10.8)
test_shapes_and_distributions_are_checkedboundarymismatched shapes and non-distributions raisevocabularies of draft and target must match
test_speculative_output_is_exactly_pstatisticaloutput distribution =p= p exactly, random p,qp, q in eighthsthe theorem L8.6 relies on
test_speculative_chi_squarestatistical20000 seeded draws match pp (chi-square)the same law on arbitrary distributions
test_speculative_accepts_without_resamplingunitan accepted draft ignores u′u'the draw order is part of the parity contract
test_speculative_rejects_bad_draftsboundaryxx outside the vocabulary or qx=0q_x = 0 raisedraft and proposal disagreeing is a bug
test_rejection_sample_first_try_enumeratedstatisticalone try accepts xx with probability px/mp_x/m, exactlythe von Neumann law
test_rejection_sample_tries_are_geometricstatisticalmean tries =m= m within 4 standard errorsthe cost of a loose envelope
test_rejection_sample_draw_orderunitproposal uniform first, then acceptance; two per tryreproducible across languages
test_rejection_sample_checks_the_envelopeboundarypx>mqxp_x > m q_x, m<1m < 1, and no acceptance in max_tries raisea wrong envelope silently biases
PitfallSymptomCaught by
1. accepting on u≤ru \le ra token with px=0p_x = 0 is accepted when u=0u = 0test_accept_is_strict (mutant s01)
2. the ratio upside down, or “capped” by inverting itover-proposed tokens kept, under-proposed ones droppedtest_accept_always_when_target_dominates (mutants s02, s12)
3. residual max⁡(0,q−p)\max(0, q - p) or ∣p−q∣\lvert p - q\rvertresamples the tokens the draft over-proposedtest_residual_is_a_distribution (mutants s03, s05)
4. forgetting to normalize the residualsample_categorical rejects it, or the CDF ends below 1test_residual_is_a_distribution (mutant s04)
5. on rejection, resampling from pp, from qq, or returning the draftoutput no longer pp: accepted tokens counted twicetest_speculative_output_is_exactly_p (mutants s06, s07, s13)
6. p=qp = q handled as 0/0, or returning the caller’s arrayNaN, or a later in-place edit corrupts the targettest_residual_when_p_equals_q (mutants s08, s14)
7. von Neumann without mm, or without checking itthe accepted tokens follow min⁡(p,q)\min(p, q)-ish weights, not pptest_rejection_sample_first_try_enumerated (mutant s15), test_rejection_sample_checks_the_envelope (mutant s19)
8. summing with numpy’s pairwise np.sumlast-bit differences against the Rust porttest_residual_sum_is_ascending (mutant s09)
DirectionModuleHow it uses this
BackM07.1sample_categorical draws the residual; UniformSource is the generator type
BackM07.0uniform random variables and independence, behind every probability above
ForwardL8.6speculative decoding: one speculative_step per draft position, acceptance_probability as the reported acceptance rate
ForwardL10.8the 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.

Your pieceProduction equivalentWhat it addsWhere to look
speculative_stepvLLM RejectionSamplerbatched acceptance over a whole draft tree on the GPU, recovered tokens from the residualvllm/v1/sample/rejection_sampler.py
residual_distributionHugging Face assisted generationthe same residual for model drafts (_speculative_sampling)transformers/generation/utils.py
rejection_sampleNumPy’s Gamma and binomial samplersrejection with tight squeeze functions to avoid evaluating the targetnumpy/random/src/distributions/distributions.c