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)
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 anddraft_probs_fullcarries the distribution it drew from (vocab_sizerequired), enabling the classicmin(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_phasewithrelaxed_topkandrelaxed_delta) buys a higher acceptance rate inside<think>regions, where exact token identity matters less than throughput. It requiresdraft_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]
- first_rejected_idx: Index of first rejected draft position
-
Return type:
-
Tuple of
(first_rejected_idx, recovered_tokens, bonus_tokens) -
Parameters:
-
- draft_tokens (TensorValue)
- target_logits (TensorValue)
- temperature (TensorValue)
- top_k (TensorValue)
- max_k (TensorValue)
- top_p (TensorValue)
- min_top_p (TensorValue)
- seed (TensorValue)
- in_thinking_phase (TensorValue | None)
- relaxed_topk (int | None)
- relaxed_delta (float | None)
- token_bitmasks (TensorValue | None)
- draft_proposal (Literal['argmax', 'sampled'])
- draft_probs_full (TensorValue | None)
- vocab_size (int | None)