class DFlashSpeculator(DraftModelSpeculator):
_speculator_name = "DFlash" # For logging, so we can share methods with subclasses
def __init__(self, vllm_config: VllmConfig, device: torch.device):
parallel_config = vllm_config.parallel_config
speculative_config = vllm_config.speculative_config
assert speculative_config is not None
vllm_config = copy.copy(vllm_config)
vllm_config.parallel_config = replace(
parallel_config,
prefill_context_parallel_size=1,
decode_context_parallel_size=(
parallel_config.decode_context_parallel_size
if speculative_config.draft_model_config.use_mla
else 1
),
)
super().__init__(vllm_config, device)
self.hidden_states = torch.zeros(
self.max_num_tokens, self.hidden_size, dtype=self.dtype, device=device
)
# Multimodal inputs not currently supported.
self.supports_mm_inputs = False
# Each request emits exactly (bonus + N mask) query tokens per step.
self.num_query_per_req = 1 + self.num_speculative_steps
self.parallel_drafting_token_id = get_parallel_drafting_token_id(
self.draft_model_config.hf_config
)
from vllm.model_executor.models.qwen3_dflash import dflash_has_any_non_causal
self.requires_non_causal = dflash_has_any_non_causal(
self.draft_model_config.hf_config
)
# Whether the anchor query position is itself a prediction. DFlash default uses
# the anchor as the bonus token (only mask tokens predict); DSpark samples from
# the anchor and the N-1 mask token positions. See _prepare_dflash_inputs_kernel
dflash_config = (
getattr(self.draft_model_config.hf_config, "dflash_config", None) or {}
)
if dflash_config.get("sample_from_anchor", False):
raise ValueError(
"sample_from_anchor=True is not supported for DFlash. "
"DFlash uses a fixed 1+N query layout where the anchor "
"is the bonus token."
)
self.sample_from_anchor = False
# Context positions for the K/V precompute, populated by prepare_dflash_inputs.
self.context_positions = torch.zeros(
self.max_num_tokens, dtype=torch.int64, device=device
)
# Per-mask-token sampling buffers. Flattened from (num_reqs, num_spec_tokens).
max_num_sampled_tokens = self.max_num_reqs * self.num_speculative_steps
self.sample_indices = torch.zeros(
max_num_sampled_tokens, dtype=torch.int64, device=device
)
self.sample_pos = torch.zeros(
max_num_sampled_tokens, dtype=torch.int64, device=device
)
# -1 marks an inert sampling row. CUDA graph capture can execute the
# full buffer before a real batch has populated it, so zero would make
# every padding row scatter into request slot 0.
self.sample_idx_mapping = torch.full(
(max_num_sampled_tokens,), -1, dtype=torch.int32, device=device
)
# [0, 1, ..., N-1, 0, 1, ..., N-1, ...] -> the per-token column index into
# draft_logits[req, step, :].
self.sample_col = torch.arange(
self.num_speculative_steps, dtype=torch.int32, device=device
).repeat(self.max_num_reqs)
self.query_cudagraph_manager: DFlashCudaGraphManager | None = None
self.draft_kv_cache_group_id: int = -1
@property
def attn_vllm_config(self) -> VllmConfig:
# The draft's attention differs from the target's in causality.
config = copy.copy(super().attn_vllm_config)
config.attention_config = replace(
self.vllm_config.attention_config,
use_non_causal=self.requires_non_causal,
)
return config
def init_cudagraph_manager(self, cudagraph_mode: CUDAGraphMode) -> None:
wants_full = cudagraph_mode.decode_mode() == CUDAGraphMode.FULL
supports_full = (
self.attn_cg_support.min_cg_support.value
>= AttentionCGSupport.UNIFORM_BATCH.value
)
if wants_full and not supports_full:
logger.warning(
"%s draft attention (%s) does not support full CUDA graphs; "
"running the draft eagerly.",
self._speculator_name,
self.attn_cg_support.min_cg_attn_backend,
)
# PIECEWISE cudagraphs are not supported for dflash.
if wants_full and supports_full:
cudagraph_mode = CUDAGraphMode.FULL_DECODE_ONLY
else:
cudagraph_mode = CUDAGraphMode.NONE
self.query_cudagraph_manager = DFlashCudaGraphManager(
self.vllm_config,
self.device,
cudagraph_mode,
decode_query_len=self.num_query_per_req,
)
def capture(self) -> None:
logger.info("Capturing model for %s speculator...", self._speculator_name)
# Padded sample rows must not scatter into a live request during capture.
self.sample_indices.zero_()
self.sample_pos.zero_()
self.sample_idx_mapping.fill_(-1)
# Capture must not write context K/V.
self._context_slot_mappings.fill_(PAD_SLOT_ID)
assert self.query_cudagraph_manager is not None
self.query_cudagraph_manager.capture(
self._generate_draft,
self.input_buffers,
self.block_tables,
self.attn_groups,
self.kv_cache_config,
self.max_model_len,
causal=self._group_causal,
precompute_context_kv=lambda num_reqs: self._precompute_context_kv(
0, self._num_graph_context_tokens(num_reqs)
),
progress_bar_desc=f"Capturing {self._speculator_name.lower()} CUDA graphs",
)
def load_draft_model(
self,
target_model: nn.Module,
target_attn_layer_names: set[str],
) -> nn.Module:
return load_dflash_model(target_model, self.vllm_config)
def set_attn(
self,
model_state: ModelState,
kv_cache_config: KVCacheConfig,
block_tables: BlockTables,
target_input_buffers: InputBuffers,
target_attn_groups: list[list[AttentionGroup]],
) -> None:
super().set_attn(
model_state,
kv_cache_config,
block_tables,
target_input_buffers,
target_attn_groups,
)
# FlashAttention's AOT split schedule is wrong for a windowed drafter,
# and `_get_sliding_window_configs` leaves it on or off depending on
# whether the target also runs FlashAttention. Decide it here instead.
for groups in self.attn_groups:
for group in groups:
builder = group.get_metadata_builder()
if getattr(
builder, "aot_schedule", False
) and get_kv_cache_spec_sliding_window(builder.kv_cache_spec):
# `aot_schedule` belongs to FlashAttention's builder, not
# to the base class this loop is typed against.
builder.aot_schedule = False # type: ignore[attr-defined]
self.draft_kv_cache_group_ids = [
gid for gid, g in enumerate(self.attn_groups) if g
]
assert self.draft_kv_cache_group_ids, "No draft attention groups found."
self.draft_kv_cache_group_id = self.draft_kv_cache_group_ids[0]
# Per-group context slot buffers for the precompute (one row per group).
self._context_slot_mappings = torch.zeros(
len(self.draft_kv_cache_group_ids),
self.max_num_tokens,
dtype=torch.int64,
device=self.device,
)
# Map each draft decoder layer to the index (within draft_kv_cache_group_ids)
# of the kv-cache group its cache belongs to. Models that share a single group
# leave this as None and share one context slot mapping.
self._layer_group_idx: list[int] | None = None
# Per-KV-group causal, falling back to whether the drafter is all-causal.
self._group_causal: dict[int, bool] | bool = not self.requires_non_causal
if hasattr(self.model, "get_draft_kv_cache_layer_names"):
layer_names = self.model.get_draft_kv_cache_layer_names()
name_to_gid = {
ln: gid
for gid, group in enumerate(kv_cache_config.kv_cache_groups)
for ln in group.layer_names
}
gid_to_idx = {gid: i for i, gid in enumerate(self.draft_kv_cache_group_ids)}
self._layer_group_idx = [
gid_to_idx[name_to_gid[name]] for name in layer_names
]
if hasattr(self.model, "get_draft_attn_causal"):
self._group_causal = {
name_to_gid[name]: layer_causal
for name, layer_causal in zip(
layer_names, self.model.get_draft_attn_causal()
)
}
@torch.inference_mode()
def _run_model(
self,
num_tokens: int,
attn_metadata: dict[str, Any] | None,
slot_mappings: dict[str, torch.Tensor] | None,
num_tokens_across_dp: torch.Tensor | None,
cudagraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
) -> torch.Tensor:
batch_descriptor = BatchDescriptor(num_tokens=num_tokens)
with set_forward_context(
attn_metadata,
self.vllm_config,
num_tokens=num_tokens,
cudagraph_runtime_mode=cudagraph_runtime_mode,
num_tokens_across_dp=num_tokens_across_dp,
slot_mapping=slot_mappings,
batch_descriptor=batch_descriptor,
):
last_hidden_states = self.model(
input_ids=self.input_buffers.input_ids[:num_tokens],
positions=self.input_buffers.positions[:num_tokens],
inputs_embeds=None,
)
return last_hidden_states
def _generate_draft(
self,
num_reqs: int,
num_tokens_padded: int,
attn_metadata: dict[str, Any] | None,
slot_mappings: dict[str, torch.Tensor] | None,
num_tokens_across_dp: torch.Tensor | None,
cudagraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
) -> None:
last_hidden_states = self._run_model(
num_tokens_padded,
attn_metadata,
slot_mappings,
num_tokens_across_dp,
cudagraph_runtime_mode,
)
num_sample = num_reqs * self.num_speculative_steps
sample_hidden_states = last_hidden_states[self.sample_indices[:num_sample]]
# sample_pos is the predicted token's position P. Sampling keys a draw
# by the position before the sampled token, P-1.
draft_tokens = self.sample_draft(
sample_hidden_states,
self.sample_pos[:num_sample] - 1,
self.sample_idx_mapping[:num_sample],
self.temperature,
self.seeds,
self.sample_col[:num_sample],
self.draft_logits,
)
self.draft_tokens[:num_reqs] = draft_tokens.view(
num_reqs, self.num_speculative_steps
)
def _num_graph_context_tokens(self, num_reqs: int) -> int:
# Context rows a captured draft step stores: one full verify per request.
return min(num_reqs * (self.num_speculative_steps + 1), self.max_num_tokens)
def _precompute_context_kv(
self, start: int, end: int, dummy_run: bool = False
) -> None:
if dummy_run:
context_slots: torch.Tensor | list[torch.Tensor | None] | None = None
elif self._layer_group_idx is not None:
context_slots = [
self._context_slot_mappings[gidx][start:end]
for gidx in self._layer_group_idx
]
else:
context_slots = self._context_slot_mappings[0][start:end]
self.model.precompute_and_store_context_kv(
self.hidden_states[start:end],
self.context_positions[start:end],
context_slots,
)
def prepare_context_anchor(
self, input_batch: InputBatch, num_rejected: torch.Tensor
) -> None:
"""Publish context features required by a draft's candidate head."""
@torch.inference_mode()
def propose(
self,
input_batch: InputBatch,
attn_metadata: dict[str, Any],
slot_mappings: dict[str, torch.Tensor],
# [num_tokens, hidden_size]
last_hidden_states: torch.Tensor,
# num_layers x [num_tokens, hidden_size]
aux_hidden_states: list[torch.Tensor] | None,
# [num_reqs]
num_sampled: torch.Tensor,
# [num_reqs]
num_rejected: torch.Tensor,
# [max_num_reqs]
last_sampled: torch.Tensor,
# [max_num_reqs]
next_prefill_tokens: torch.Tensor,
# [max_num_reqs]
temperature: torch.Tensor,
# [max_num_reqs]
seeds: torch.Tensor,
dp_sync: DPSyncState | None = None,
dummy_run: bool = False,
skip_attn_for_dummy_run: bool = False,
mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None,
is_profile: bool = False,
) -> torch.Tensor:
num_reqs = input_batch.num_reqs
num_target_tokens = input_batch.num_tokens
num_query_tokens = num_reqs * self.num_query_per_req
max_seq_len = input_batch.seq_lens_cpu_upper_bound[:num_reqs].max().item()
self.draft_max_seq_len = min(
max_seq_len + self.num_query_per_req, self.max_model_len
)
# NOTE: To avoid CPU-GPU synchronization without CPU knowing the
# number of rejected tokens, we maintain the size of input_ids and
# hidden_states the same as the target model's. This means, we pad each
# request's query length to include any rejected positions.
if aux_hidden_states:
hidden_states = self.model.combine_hidden_states(
torch.cat(aux_hidden_states, dim=-1)
)
else:
hidden_states = last_hidden_states
self.hidden_states[:num_target_tokens].copy_(hidden_states[:num_target_tokens])
self.prepare_context_anchor(input_batch, num_rejected)
if dummy_run and skip_attn_for_dummy_run:
# Memory profiling path: block_tables / kv_cache_config are not initialized.
# Since DFlash needs to build its own attention metadata, we must skip the
# preparation in this path and run a minimal forward pass.
self.model.precompute_and_store_context_kv(
self.hidden_states[:num_target_tokens],
self.context_positions[:num_target_tokens],
)
# DFlash processes all speculative tokens in one forward pass,
# so the real token count is num_query_tokens.
self._prepare_eplb_forward(num_query_tokens)
self._generate_draft(
num_reqs,
num_query_tokens,
attn_metadata=None,
slot_mappings=None,
num_tokens_across_dp=None,
cudagraph_runtime_mode=CUDAGraphMode.NONE,
)
return self.draft_tokens[:num_reqs]
if self.pcp_manager is not None and not dummy_run:
self.block_tables.gather_block_tables(
input_batch.idx_mapping, num_reqs_padded=num_reqs
)
# The query slot mapping is written into the shared BlockTables slot_mappings.
# That buffer's address is what the captured CUDA graph reads from at replay.
assert self.draft_kv_cache_group_id >= 0
# Support multiple draft KV cache groups by preparing inputs once for each
for i, gid in enumerate(self.draft_kv_cache_group_ids):
prepare_dflash_inputs(
self.input_buffers,
self.block_tables.slot_mappings[gid],
self.context_positions,
self._context_slot_mappings[i],
self.sample_indices,
self.sample_pos,
self.sample_idx_mapping,
self.temperature,
self.seeds,
input_batch,
num_sampled,
num_rejected,
last_sampled,
next_prefill_tokens,
temperature,
seeds,
self.block_tables.input_block_tables[gid],
self.block_tables.kernel_block_sizes[gid],
self.block_tables.cp_rank,
self.dcp_size,
self.block_tables.cp_interleave,
self.parallel_drafting_token_id,
self.num_query_per_req,
self.num_speculative_steps,
self.max_num_reqs,
self.max_num_tokens,
self.max_model_len,
self.sample_from_anchor,
)
batch_sync, num_batch_tokens = (
self._build_uniform_batch_dp_sync(dp_sync, num_reqs, self.num_query_per_req)
if dp_sync is not None
else (None, num_query_tokens)
)
# Every DFlash step has exactly num_query_per_req tokens, so we can use FULL CGs
batch_desc, batch_sync = dispatch_cg_and_sync_dp(
self.query_cudagraph_manager,
num_reqs,
num_batch_tokens,
uniform_token_count=self.num_query_per_req,
dp_size=self.dp_size,
dp_rank=self.dp_rank,
need_eager=is_profile,
dp_sync=batch_sync,
)
num_tokens_padded = batch_desc.num_tokens
num_tokens_across_dp = (
batch_sync.num_tokens_across_dp if batch_sync is not None else None
)
if batch_desc.cg_mode == CUDAGraphMode.FULL:
# The graph stores the first num_context context rows.
assert batch_desc.num_reqs is not None
num_context = self._num_graph_context_tokens(batch_desc.num_reqs)
if dummy_run:
# Dummy block tables are placeholders: write no context K/V.
self._context_slot_mappings[:, :num_context].fill_(PAD_SLOT_ID)
elif num_target_tokens <= num_context:
# Rows past the batch keep stale positions but write no K/V.
self._context_slot_mappings[:, num_target_tokens:num_context].fill_(
PAD_SLOT_ID
)
else:
# Prefill context beyond the graph's rows is stored before replay.
self._precompute_context_kv(num_context, num_target_tokens)
else:
self._precompute_context_kv(0, num_target_tokens, dummy_run)
# Rebuild the draft attention metadata even when replaying the FULL
# graph so that any attention metadata builder state is updated.
draft_attn_metadata = self._build_uniform_attn_metadata(
num_reqs=num_reqs,
batch_desc=batch_desc,
num_query_per_req=self.num_query_per_req,
seq_lens_cpu_upper_bound=input_batch.seq_lens_cpu_upper_bound,
step=self.num_query_per_req,
causal=self._group_causal,
)
draft_slot_mappings_by_layer = build_slot_mappings_by_layer(
self.block_tables.slot_mappings[:, :num_tokens_padded],
self.kv_cache_config,
)
# DFlash processes all speculative tokens in one forward pass,
# so the real token count is num_query_tokens.
self._prepare_eplb_forward(num_query_tokens)
if batch_desc.cg_mode == CUDAGraphMode.FULL:
assert self.query_cudagraph_manager is not None
self.query_cudagraph_manager.run_fullgraph(batch_desc)
else:
self._generate_draft(
num_reqs,
num_tokens_padded,
draft_attn_metadata,
draft_slot_mappings_by_layer,
num_tokens_across_dp=num_tokens_across_dp,
cudagraph_runtime_mode=batch_desc.cg_mode,
)
return self.draft_tokens[:num_reqs]