class PlainDraftModelSpeculator(DraftModelSpeculator):
"""Speculative decoding using a separate smaller draft LM.
Unlike Eagle, the draft model runs fully independently of the target model.
Step 0 builds an expanded buffer (accepted + correction token + rejected
slots) via a Triton kernel; steps 1..k-1 are single-token decode steps.
"""
# Plain draft model needs one extra slot per request for correction.
num_extra_query_per_req = 1
def __init__(self, vllm_config: VllmConfig, device: torch.device):
super().__init__(vllm_config, device)
# draft_max_seq_len is read by the parent's attention metadata builder.
# Plain draft model doesn't do per-batch adjustment; cap at max.
self.draft_max_seq_len = self.max_model_len
self.last_token_indices = torch.zeros(
self.max_num_reqs, dtype=torch.int64, device=device
)
# GPU-side arange for decode query_start_loc initialisation.
self.arange_gpu = torch.arange(
self.max_num_reqs + 1, dtype=torch.int32, device=device
)
# Scalar step counter used as column index into draft_logits.
self.current_draft_step = torch.tensor(0, dtype=torch.int64, device=device)
_expanded_max = self.max_num_tokens + self.max_num_reqs
self.expanded_input_ids = torch.zeros(
_expanded_max, dtype=torch.int32, device=device
)
self.expanded_positions = torch.zeros(
_expanded_max, dtype=torch.int64, device=device
)
self.expanded_slot_mappings: torch.Tensor
self.supports_mm_inputs = False
def init_cudagraph_manager(self, cudagraph_mode: CUDAGraphMode) -> None:
pass # CUDA graph not yet supported for plain draft model speculator.
def capture(self) -> None:
pass # CUDA graph not yet supported for plain draft model speculator.
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,
)
self.expanded_slot_mappings = torch.empty(
block_tables.num_kv_cache_groups,
self.max_num_tokens + self.max_num_reqs,
dtype=torch.int64,
device=self.device,
)
def load_draft_model(
self,
target_model: nn.Module,
target_attn_layer_names: set[str],
) -> nn.Module:
spec = self.speculative_config
assert spec is not None
draft_vllm_config = replace(
self.vllm_config,
model_config=self.draft_model_config,
quant_config=None,
parallel_config=replace(
spec.draft_parallel_config,
rank=self.vllm_config.parallel_config.rank,
),
)
with set_model_tag("draft_model"):
return get_model(
vllm_config=draft_vllm_config,
prefix="draft_model",
)
@torch.inference_mode()
def _run_model(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
attn_metadata: dict[str, Any] | None,
slot_mappings: dict[str, torch.Tensor] | None,
num_tokens_across_dp: torch.Tensor | None,
) -> torch.Tensor:
num_tokens = input_ids.shape[0]
batch_descriptor = BatchDescriptor(num_tokens=num_tokens)
with set_forward_context(
attn_metadata,
self.vllm_config,
num_tokens=num_tokens,
cudagraph_runtime_mode=CUDAGraphMode.NONE,
num_tokens_across_dp=num_tokens_across_dp,
slot_mapping=slot_mappings,
batch_descriptor=batch_descriptor,
):
hidden_states = self.model( # type: ignore[misc]
input_ids=input_ids,
positions=positions,
)
if isinstance(hidden_states, tuple):
hidden_states = hidden_states[0]
return hidden_states
def _accepted_last_indices(
self,
input_batch: InputBatch,
num_rejected: torch.Tensor,
num_reqs: int,
) -> torch.Tensor:
qsl = input_batch.query_start_loc
adjusted_lens = qsl[1 : num_reqs + 1] - qsl[:num_reqs] - num_rejected[:num_reqs]
return qsl[:num_reqs] + adjusted_lens - 1
def _prefill(
self,
input_batch: InputBatch,
num_rejected: torch.Tensor,
last_sampled: torch.Tensor,
num_reqs: int,
skip_attn: bool,
num_tokens_across_dp: torch.Tensor | None,
) -> None:
if skip_attn:
src_tokens = input_batch.num_tokens
self.last_token_indices[:num_reqs] = self._accepted_last_indices(
input_batch, num_rejected, num_reqs
)
hidden_states = self._run_model(
input_batch.input_ids[:src_tokens],
input_batch.positions[:src_tokens],
None,
None,
num_tokens_across_dp,
)
positions_buf = input_batch.positions
else:
block_tables = self.block_tables
kv_cache_config = self.kv_cache_config
assert block_tables is not None
assert kv_cache_config is not None
total_expanded = prepare_prefill_inputs(
input_buffers=self.input_buffers,
input_batch=input_batch,
last_sampled=last_sampled,
num_rejected=num_rejected,
expanded_input_ids=self.expanded_input_ids,
expanded_positions=self.expanded_positions,
last_token_indices=self.last_token_indices,
max_num_reqs=self.max_num_reqs,
max_model_len=self.max_model_len,
)
prefill_slot_mappings = block_tables.compute_slot_mappings(
input_batch.idx_mapping,
self.input_buffers.query_start_loc[: num_reqs + 1],
self.expanded_positions[:total_expanded],
total_expanded,
out=self.expanded_slot_mappings,
)
qsl_np = input_batch.query_start_loc_np
query_start_loc_cpu_expanded = (
torch.from_numpy(qsl_np[: num_reqs + 1]).int()
+ torch.from_numpy(self.arange_np[: num_reqs + 1]).int()
)
prefill_attn_md = self._build_attn_metadata(
batch_desc=BatchExecutionDescriptor(
cg_mode=CUDAGraphMode.NONE,
num_tokens=total_expanded,
num_reqs=num_reqs,
),
num_reqs=num_reqs,
query_start_loc_np=query_start_loc_cpu_expanded.numpy(),
seq_lens_cpu_upper_bound=input_batch.seq_lens_cpu_upper_bound,
step=0,
slot_mappings=prefill_slot_mappings,
)
prefill_slot_maps_by_layer = build_slot_mappings_by_layer(
prefill_slot_mappings, kv_cache_config
)
hidden_states = self._run_model(
self.expanded_input_ids[:total_expanded],
self.expanded_positions[:total_expanded],
prefill_attn_md,
prefill_slot_maps_by_layer,
num_tokens_across_dp,
)
positions_buf = self.expanded_positions
last_indices = self.last_token_indices[:num_reqs]
self.current_draft_step.fill_(0)
self.draft_tokens[:num_reqs, 0] = self.sample_draft(
hidden_states[last_indices],
positions_buf[last_indices],
self.idx_mapping[:num_reqs],
self.temperature,
self.seeds,
self.current_draft_step,
self.draft_logits,
)
def _multi_step_decode(
self,
input_batch: InputBatch,
num_rejected: torch.Tensor,
num_reqs: int,
skip_attn: bool,
num_tokens_across_dp: torch.Tensor | None,
) -> None:
if skip_attn:
last_positions = input_batch.positions[
self._accepted_last_indices(input_batch, num_rejected, num_reqs)
]
else:
last_positions = self.expanded_positions[self.last_token_indices[:num_reqs]]
self.input_buffers.positions[:num_reqs].copy_(last_positions)
# The decode loop increments seq_lens BEFORE each forward.
# Initial seq_lens = target_seq_lens - num_rejected + 1.
# Step 1 forward uses target_seq_lens - num_rejected + 2.
# Step 2 forward uses target_seq_lens - num_rejected + 3.
self.input_buffers.seq_lens[:num_reqs].copy_(
torch.clamp(
input_batch.seq_lens[:num_reqs] - num_rejected[:num_reqs].int() + 1,
max=self.max_model_len,
)
)
self.input_buffers.query_start_loc[: num_reqs + 1].copy_(
self.arange_gpu[: num_reqs + 1]
)
idx_mapping = input_batch.idx_mapping
for step in range(1, self.num_speculative_steps):
self.input_buffers.input_ids[:num_reqs].copy_(
self.draft_tokens[:num_reqs, step - 1].int()
)
torch.clamp(
self.input_buffers.positions[:num_reqs] + 1,
max=self.max_model_len - 1,
out=self.input_buffers.positions[:num_reqs],
)
# Increment seq_lens BEFORE the forward so the attention reads from
# the KV slot written by the previous step.
torch.clamp(
self.input_buffers.seq_lens[:num_reqs] + 1,
max=self.max_model_len,
out=self.input_buffers.seq_lens[:num_reqs],
)
decode_attn_md = None
step_slot_maps_by_layer = None
if not skip_attn:
block_tables = self.block_tables
kv_cache_config = self.kv_cache_config
assert block_tables is not None
assert kv_cache_config is not None
q_start = self.input_buffers.query_start_loc[: num_reqs + 1]
positions = self.input_buffers.positions[:num_reqs]
step_slot_mappings = block_tables.compute_slot_mappings(
idx_mapping, q_start, positions, num_reqs
)
step_slot_maps_by_layer = build_slot_mappings_by_layer(
step_slot_mappings, kv_cache_config
)
decode_attn_md = self._build_uniform_attn_metadata(
batch_desc=BatchExecutionDescriptor(
cg_mode=CUDAGraphMode.NONE,
num_tokens=num_reqs,
num_reqs=num_reqs,
uniform_token_count=1,
),
num_reqs=num_reqs,
num_query_per_req=1,
seq_lens_cpu_upper_bound=input_batch.seq_lens_cpu_upper_bound,
# Include the correction token inserted during prefill.
step=step + 1,
)
hidden_states = self._run_model(
self.input_buffers.input_ids[:num_reqs],
self.input_buffers.positions[:num_reqs],
decode_attn_md,
step_slot_maps_by_layer,
num_tokens_across_dp,
)
self.current_draft_step.fill_(step)
self.draft_tokens[:num_reqs, step] = self.sample_draft(
hidden_states[:num_reqs],
self.input_buffers.positions[:num_reqs],
self.idx_mapping[:num_reqs],
self.temperature,
self.seeds,
self.current_draft_step,
self.draft_logits,
)
@torch.inference_mode()
def propose(
self,
input_batch: InputBatch,
attn_metadata: dict[str, Any],
slot_mappings: dict[str, torch.Tensor],
last_hidden_states: torch.Tensor, # unused
aux_hidden_states: list[torch.Tensor] | None, # unused
num_sampled: torch.Tensor,
num_rejected: torch.Tensor,
last_sampled: torch.Tensor,
next_prefill_tokens: torch.Tensor,
temperature: torch.Tensor,
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:
assert self.model is not None
num_reqs = input_batch.num_reqs
skip_attn = dummy_run and skip_attn_for_dummy_run
num_tokens_across_dp = (
dp_sync.num_tokens_across_dp if dp_sync is not None else None
)
# Copy per-request temperature/seeds/idx_mapping into pre-allocated
# buffers so sample_draft can read them without extra slicing.
self._copy_request_inputs(
num_reqs,
input_batch.idx_mapping,
temperature,
seeds,
)
self._prefill(
input_batch,
num_rejected,
last_sampled,
num_reqs,
skip_attn,
num_tokens_across_dp,
)
if self.num_speculative_steps == 1:
return self.draft_tokens[:num_reqs, :1]
self._multi_step_decode(
input_batch,
num_rejected,
num_reqs,
skip_attn,
num_tokens_across_dp,
)
return self.draft_tokens[:num_reqs]