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, return_n_logits, host_input_row_offsets, data_parallel_splits, batch_context_lengths, ep_inputs, devices, data_parallel_degree, merged_tokens, merged_offsets, host_merged_offsets, merged_offsets_per_dev)
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]])
- return_n_logits (TensorValue)
- host_input_row_offsets (TensorValue)
- data_parallel_splits (TensorValue)
- batch_context_lengths (list[TensorValue])
- ep_inputs (list[Value[Any]] | None)
- devices (Sequence[DeviceRef])
- data_parallel_degree (int)
- merged_tokens (TensorValue)
- merged_offsets (TensorValue)
- host_merged_offsets (TensorValue)
- merged_offsets_per_dev (list[TensorValue])
batch_context_lengths
batch_context_lengths: list[TensorValue]
data_parallel_degree
data_parallel_degree: int
data_parallel_splits
data_parallel_splits: TensorValue
device0
property device0: DeviceRef
The device that owns the batch-wide (non-sharded) tensors.
devices
draft_kv_collections
draft_kv_collections: list[KVCacheInputsPerDevice[TensorValue, BufferValue]]
draft_tokens
draft_tokens: TensorValue
ep_inputs
host_input_row_offsets
host_input_row_offsets: TensorValue
host_merged_offsets
host_merged_offsets: TensorValue
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]
merged_tokens
merged_tokens: TensorValue
n_devs
property n_devs: int
Number of devices the target and draft are sharded across.
num_draft_tokens
property num_draft_tokens: Dim
How many tokens the previous iteration proposed, K.
return_n_logits
return_n_logits: TensorValue
signal_buffers
signal_buffers: list[BufferValue]
tokens
tokens: TensorValue