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 class
SequentialProposer
SequentialProposer
class max.pipelines.speculative.driver.SequentialProposer(*args, **kwargs)
Bases: Protocol[_TargetHiddenT]
A draft that emits one token per step, K steps deep.
carry_dim_names
carry_dim_names: CarryDimNames
How each step’s carry dims are named.
decode_swaps
decode_swaps: tuple[DecodeKVSwap, ...]
Which draft-cache fields to retarget to q = 1 before the loop.
draft_cache
draft_cache: DraftCache
Which cache the draft writes, and so what the driver advances.
hidden_dim
Trailing dim of the carried hidden state, used when rebinding it.
passthrough_decode_swaps
Which SequentialBatch.passthrough_kv leaves swap too.
The primary draft leaf always takes decode_swaps. A paired cache
is only sometimes symmetric.
prefill()
prefill(batch, tokens, target_hidden)
Runs draft step 0 over the whole target-corrected sequence.
-
Parameters:
-
- batch (SequentialBatch)
- tokens (TensorValue)
- target_hidden (_TargetHiddenT)
-
Return type:
reuse
reuse: ReuseSpec | None
The result carried beside the hidden state, or None for a draft
with none.
split_prefix
split_prefix: str
Names the accepted-position gather and each step’s draft subgraph.
step()
step(batch, draft_input, index)
Runs draft step index over one token per batch element.
-
Parameters:
-
- batch (SequentialBatch)
- draft_input (DraftStepInput)
- index (int)
-
Return type:
step_hidden_mode
step_hidden_mode: ReturnHiddenStates
LAST_PER_DEVICE returns post-allgather full-batch tensors and so
needs a DP slice; ALL returns per-device ones and must not be sliced.
The driver applies that rule once.
uses_thinking_phase
uses_thinking_phase: bool
Whether the draft’s acceptance test reads in_thinking_phase.