IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).

Python function

stochastic_acceptance_sampler

stochastic_acceptance_sampler()​

max.nn.sampling.stochastic_acceptance_sampler(draft_tokens, target_logits, temperature, top_k, max_k, top_p, min_top_p, seed, in_thinking_phase=None, relaxed_topk=None, relaxed_delta=None, token_bitmasks=None, draft_proposal='argmax', draft_probs_full=None, vocab_size=None)

source

Verifies speculative draft tokens against the target model.

Speculative decoding is only lossless if every committed token is distributed exactly as if the target model had sampled it alone, under the request’s own sampling params (temperature, top_k/top_p/min_top_p). This graph enforces that invariant: it decides how many draft tokens to accept, replaces the first rejected position with a recovered token, and produces the bonus token that is committed when every draft was accepted.

How acceptance is judged depends on how the draft chose its tokens, because the correct math differs:

  • draft_proposal="argmax" (the default): the draft proposed deterministically, so acceptance reduces to matching a sample drawn from the truncated target distribution. See _argmax_draft_verdict().
  • draft_proposal="sampled": the draft sampled stochastically and draft_probs_full carries the distribution it drew from (vocab_size required), enabling the classic min(1, p / q) ratio test with residual recovery. See _sampled_draft_verdict().

Two per-row overlays deliberately trade the losslessness guarantee for other goals:

  • Relaxed thinking acceptance (in_thinking_phase with relaxed_topk and relaxed_delta) buys a higher acceptance rate inside <think> regions, where exact token identity matters less than throughput. It requires draft_proposal="argmax" and raises otherwise. See _relaxed_thinking_verdict().
  • Rows with ~zero temperature use greedy verification: the draft is accepted iff it equals the target argmax, which is also the recovered and bonus token.

When token_bitmasks is provided, grammar constraints mask the target logits before anything is sampled, so recovered and bonus tokens always satisfy structured-output constraints.

Returns:

  • first_rejected_idx: Index of first rejected draft position [batch]
  • recovered_tokens: Tokens sampled from target distribution [batch, num_steps]
  • bonus_tokens: Bonus token from final position [batch, 1]

Return type:

Tuple of (first_rejected_idx, recovered_tokens, bonus_tokens)

Parameters: