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 module
max.pipelines.architectures.unified_dflash_llama3
DFlash speculative decoding for Llama3 with unified graph compilation.
DflashDraftHFConfig
class max.pipelines.architectures.unified_dflash_llama3.DflashDraftHFConfig(mask_token_id, target_layer_ids, block_size=None, num_target_layers=None)
Bases: object
Parsed DFlash fields from a draft HuggingFace config.
-
Parameters:
block_size
draft_width()
draft_width(speculative, *, warn=True)
Returns the draft width, which is block_size - 1.
The drafter only works at its trained block size, so a width that
disagrees is replaced with a warning. A checkpoint with no
block_size needs an explicit width.
-
Parameters:
-
- speculative (SpeculativeConfig)
- warn (bool)
-
Return type:
mask_token_id
mask_token_id: int
num_target_layers
target_layer_ids
PersistentInputBuffers
class max.pipelines.architectures.unified_dflash_llama3.PersistentInputBuffers(tokens, input_row_offsets)
Bases: object
Pinned-host buffers reused across unified spec-decode batch steps.
alloc()
classmethod alloc(max_batch_size, max_batch_input_tokens, device)
Allocates persistent token and row-offset buffers for spec-decode batching.
-
Parameters:
-
Return type:
input_row_offsets
input_row_offsets: Buffer
tokens
tokens: Buffer
UnifiedDflashLlama3Config
class max.pipelines.architectures.unified_dflash_llama3.UnifiedDflashLlama3Config(*, target: 'Llama3Config', draft: 'Llama3Config', speculative_config: 'SpeculativeConfig', target_layer_ids: 'list[int]' = <factory>, mask_token_id: 'int' = 0, block_size: 'int' = 0, quantization_encoding: 'SupportedEncoding | None' = None, resolved_num_speculative_tokens: 'int | None' = None)
Bases: ArchConfigWithKVCache
-
Parameters:
-
- target (Llama3Config)
- draft (Llama3Config)
- speculative_config (SpeculativeConfig)
- target_layer_ids (list[int])
- mask_token_id (int)
- block_size (int)
- quantization_encoding (SupportedEncoding | None)
- resolved_num_speculative_tokens (int | None)
DEFAULT_ENCODING
DEFAULT_ENCODING: ClassVar[SupportedEncoding] = 'bfloat16'
SUPPORTED_ENCODINGS
SUPPORTED_ENCODINGS: ClassVar[set[SupportedEncoding]] = {'bfloat16', 'float32'}
block_size
block_size: int = 0
calculate_max_seq_len()
classmethod calculate_max_seq_len(huggingface_config, model_config)
Returns the resolved maximum sequence length.
Bounds or defaults the user’s model_config.max_length with the
model’s own limits. Construction runs this once and stores the
result on model_config.max_length; memory planning may lower
it further, but only on the memory plan.
-
Parameters:
-
- huggingface_config (AutoConfig) – The HuggingFace config to read model bounds from.
- model_config (MAXModelConfig) – The model config whose
max_lengthcarries the user’s setting.
-
Return type:
draft
draft: Llama3Config
get_kv_params()
get_kv_params()
KV cache parameters to use when running the model.
-
Return type:
get_max_seq_len()
get_max_seq_len()
Returns the effective maximum sequence length for the model.
For configs that store a deployment length, this is the value
initialize received; for metadata-only configs it derives from
the checkpoint.
-
Return type:
initialize()
classmethod initialize(pipeline_config, model_config=None, *, max_seq_len)
Initialize the config from a PipelineConfig.
-
Parameters:
-
- pipeline_config (PipelineConfig) – The pipeline configuration.
- model_config (MAXModelConfig | None) – The model configuration to read from. When
None(the default),pipeline_config.modelis used. Pass an explicit config (e.g.pipeline_config.draft_model) to initialize the arch config for a different model. - max_seq_len (int) – The effective maximum sequence length to store on
the config. The value is received, never derived here: the
pipeline model passes the memory plan’s VRAM-clamped length,
while memory planning (which runs before a plan exists)
passes the construction-resolved
model_config.max_length. Configs whose sequence length is pure model metadata (e.g. diffusion components) ignore it.
-
Return type:
mask_token_id
mask_token_id: int = 0
quantization_encoding
quantization_encoding: SupportedEncoding | None = None
resolve_block_size()
resolve_block_size(*, default=None)
resolved_num_speculative_tokens
explicit value if set, else the trained width.
-
Type:
-
Per-step draft count
speculative_config
speculative_config: SpeculativeConfig
target
target: Llama3Config
target_layer_ids
validate_dflash_fields()
validate_dflash_fields()
Strict validation run from UnifiedDflashLlama3Model.load_model
once the DFlash-specific fields have been populated from the draft
HF config — __post_init__ accepts the empty-placeholder config
produced by initialize() so we can’t enforce these there.
-
Return type:
-
None
UnifiedDflashLlama3Inputs
class max.pipelines.architectures.unified_dflash_llama3.UnifiedDflashLlama3Inputs(tokens, input_row_offsets, return_n_logits, *, kv_cache_inputs=None, lora_buffers=(), vision_embeddings=<factory>, vision_scatter_indices=<factory>, hidden_states=None, draft_tokens=None, draft_probs_full=None, seed=None, temperature=None, top_k=None, max_k=None, top_p=None, min_top_p=None, in_thinking_phase=None, pinned_bitmask=None, wait_payload=None, device_bitmask_scratch=None, structured_output=False, sampled_draft_proposal=False)
Bases: UnifiedSpecDecodeInputs
Inputs for the unified DFlash Llama3 graph.
The spec-decode fields and trailing buffer packing come from
UnifiedSpecDecodeInputs; tokens / input_row_offsets /
return_n_logits plus the KV cache form this single-device graph’s
prefix. The DFlash graph does not bind in_thinking_phase.
-
Parameters:
-
- tokens (Buffer)
- input_row_offsets (Buffer)
- return_n_logits (Buffer)
- kv_cache_inputs (KVCacheInputsInterface[Buffer, Buffer] | None)
- lora_buffers (tuple[Buffer, ...])
- vision_embeddings (list[Buffer])
- vision_scatter_indices (list[Buffer])
- hidden_states (Buffer | list[Buffer] | None)
- draft_tokens (Buffer | None)
- draft_probs_full (Buffer | None)
- seed (Buffer | None)
- temperature (Buffer | None)
- top_k (Buffer | None)
- max_k (Buffer | None)
- top_p (Buffer | None)
- min_top_p (Buffer | None)
- in_thinking_phase (Buffer | None)
- pinned_bitmask (Buffer | None)
- wait_payload (Buffer | None)
- device_bitmask_scratch (Buffer | None)
- structured_output (bool)
- sampled_draft_proposal (bool)
buffers
Returns positional Buffer inputs for model ABI calls.
input_row_offsets
input_row_offsets: Buffer
return_n_logits
return_n_logits: Buffer
tokens
tokens: Buffer
UnifiedDflashLlama3Model
class max.pipelines.architectures.unified_dflash_llama3.UnifiedDflashLlama3Model(pipeline_config, session, devices, kv_cache_config, weights, *, memory_plan, adapter=None, return_logits=ReturnLogits.LAST_TOKEN, return_hidden_states=ReturnHiddenStates.NONE, max_batch_size=1)
Bases: _UnifiedSpecDecodeModelMixin, GraphPipelineModelWithKVCache[TextContext]
Unified DFlash Llama3: target + draft in one compiled graph.
-
Parameters:
-
- pipeline_config (PipelineConfig)
- session (InferenceSession)
- devices (list[Device])
- kv_cache_config (KVCacheConfig)
- weights (Weights)
- memory_plan (MemoryPlan)
- adapter (WeightsAdapter | None)
- return_logits (ReturnLogits)
- return_hidden_states (ReturnHiddenStates)
- max_batch_size (int)
batch_processor_cls
batch_processor_cls
alias of UnifiedDflashLlama3BatchProcessor
get_kv_params()
classmethod get_kv_params(huggingface_config, pipeline_config, devices, kv_cache_config, cache_dtype)
Returns the KV cache params for the pipeline model.
Delegates to model_config_cls.construct_kv_params(...).
Subclasses with custom KV behavior should override this method.
-
Parameters:
-
- huggingface_config (PreTrainedConfig)
- pipeline_config (PipelineConfig)
- devices (list[DeviceRef])
- kv_cache_config (KVCacheConfig)
- cache_dtype (DType)
-
Return type:
model
model: Model
model_config_cls
model_config_cls
alias of UnifiedDflashLlama3Config