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
SequentialBatch
SequentialBatch
class max.pipelines.speculative.driver.SequentialBatch(tokens, input_row_offsets, draft_tokens, signal_buffers, kv_collections, draft_kv_collections, passthrough_kv, draft_cache_lengths, return_n_logits, distributed, batch_context_lengths, ep_inputs, devices, data_parallel_degree, vision_embeddings, vision_scatter_indices, merged_tokens, merged_offsets, merged_offsets_per_dev, query_offsets_per_dev, extra, num_accepted)
Bases: object
One spec-decode iteration’s graph inputs, merged and broadcast.
Built by the driver in phase 1 so that the target adapter, the proposer and the per-step loop all read the same merged offsets rather than each recomputing the broadcast.
-
Parameters:
-
- tokens (TensorValue)
- input_row_offsets (TensorValue)
- draft_tokens (TensorValue)
- signal_buffers (list[BufferValue])
- kv_collections (list[KVCacheInputsPerDevice[TensorValue, BufferValue]])
- draft_kv_collections (list[KVCacheInputsPerDevice[TensorValue, BufferValue]])
- passthrough_kv (Mapping[str, list[KVCacheInputsPerDevice[TensorValue, BufferValue]]])
- draft_cache_lengths (list[TensorValue])
- return_n_logits (TensorValue)
- distributed (DistributedInputs | None)
- batch_context_lengths (list[TensorValue])
- ep_inputs (list[Value[Any]] | None)
- devices (Sequence[DeviceRef])
- data_parallel_degree (int)
- vision_embeddings (list[TensorValue])
- vision_scatter_indices (list[TensorValue])
- merged_tokens (TensorValue)
- merged_offsets (TensorValue)
- merged_offsets_per_dev (list[TensorValue])
- query_offsets_per_dev (list[TensorValue])
- extra (Mapping[str, Any])
- num_accepted (TensorValue | None)
batch_context_lengths
batch_context_lengths: list[TensorValue]
data_parallel_degree
data_parallel_degree: int
device0
property device0: DeviceRef
The device that owns the batch-wide (non-sharded) tensors.
devices
dist
property dist: DistributedInputs
distributed, for a model that requires it.
A sharded target or draft cannot run without the host mirrors and the splits, so reaching for them on a single-device graph is a wiring mistake rather than something to fall back from.
distributed
distributed: DistributedInputs | None
The distributed-only inputs, None on a single-device graph.
draft_cache_lengths
draft_cache_lengths: list[TensorValue]
Per-device draft cache lengths for this step.
Always equal to each draft_kv_collections entry’s
cache_lengths – except for a draft that reads a cache it does not own
(SequentialProposer.draft_cache is TARGET), where the count is
a RoPE position rather than a write pointer and the collections keep the
lengths they came in with.
draft_kv_collections
draft_kv_collections: list[KVCacheInputsPerDevice[TensorValue, BufferValue]]
draft_tokens
draft_tokens: TensorValue
ep_inputs
extra
Graph inputs the driver carries but never reads, keyed by the model.
A model whose signature declares inputs outside the canonical set reaches them from its adapters through here, rather than the driver growing a field per model for values none of its phases understand.
input_row_offsets
input_row_offsets: TensorValue
kv_collections
kv_collections: list[KVCacheInputsPerDevice[TensorValue, BufferValue]]
merged_offsets
merged_offsets: TensorValue
merged_offsets_per_dev
merged_offsets_per_dev: list[TensorValue]
The verify window’s offsets, on every device. Stable across the loop.
merged_tokens
merged_tokens: TensorValue
n_devs
property n_devs: int
Number of devices the target and draft are sharded across.
num_accepted
num_accepted: TensorValue | None
How many draft tokens each request accepted, None before the accept.
Set for the whole propose phase, so a draft can derive per-step state from the position it continues from, without the driver knowing what that state is.
num_draft_tokens
property num_draft_tokens: Dim
How many tokens the previous iteration proposed, K.
passthrough_kv
passthrough_kv: Mapping[str, list[KVCacheInputsPerDevice[TensorValue, BufferValue]]]
Cache leaves beyond the primary pair, keyed by the model’s name for them.
The driver drives exactly one target leaf and one draft leaf: it advances
the draft leaf’s cache lengths and applies the declared decode swaps to it.
A model whose attention is split across paired caches names the leaf the
driver should drive as the primary one and reaches the rest through here,
listing in SequentialProposer.passthrough_decode_swaps any that
need the same q = 1 retarget.
query_offsets_per_dev
query_offsets_per_dev: list[TensorValue]
Offsets over this draft call’s query, on every device.
The merged offsets for step 0, which runs over the whole corrected
sequence; a one-token-per-request ramp for steps 1..K-1. A draft that
cross-attends into the target’s cache needs both these and the stable
merged_offsets_per_dev, which is why they are separate fields.
return_n_logits
return_n_logits: TensorValue
signal_buffers
signal_buffers: list[BufferValue]
tokens
tokens: TensorValue
vision_embeddings
vision_embeddings: list[TensorValue]
Per-device merged vision embeddings, empty for a text-only target.
A vision target scatters these into the merged sequence before running its stack, so they reach the target adapter rather than the driver’s phases – which never read them.
vision_scatter_indices
vision_scatter_indices: list[TensorValue]
Per-device merge positions for vision_embeddings.