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

synthetic_acceptance_sampler

synthetic_acceptance_sampler()​

max.nn.sampling.synthetic_acceptance_sampler(draft_tokens, target_logits, base_acceptance_rate, num_draft_steps, seed, *, temperature=None, top_k=None, max_k=None, top_p=None, min_top_p=None)

source

Synthetic sampler for speculative decoding benchmarking.

Accepts each draft position independently with probability base_acceptance_rate. Once a position is rejected all subsequent positions are also rejected. Accepted positions commit the draft token, so generated text is not a faithful speculative decode.

With sampling params, recovered and bonus tokens are the draws of stochastic_acceptance_sampler(), whose position 0 uses the same kernel and seed as plain decoding: at rate 0 the output matches decoding without speculation. Without them, they are the target argmax.

Parameters:

  • draft_tokens (TensorValue) – Draft token ids [batch, num_steps].
  • target_logits (TensorValue) – Verified target logits.
  • base_acceptance_rate (float) – Per-position acceptance probability.
  • num_draft_steps (int) – Number of speculative draft steps.
  • seed (TensorValue) – Per-execute seed tensor. The accept draws use its row-0 value.
  • temperature (TensorValue | None) – Per-row sampling params, all or none.
  • top_k (TensorValue | None) – Per-row sampling params, all or none.
  • max_k (TensorValue | None) – Per-row sampling params, all or none.
  • top_p (TensorValue | None) – Per-row sampling params, all or none.
  • min_top_p (TensorValue | None) – Per-row sampling params, all or none.

Return type:

tuple[TensorValue, TensorValue, TensorValue]

Returns (first_rejected_idx, recovered_tokens, bonus_tokens)