class DSparkSpeculator(DFlashSpeculator):
_speculator_name = "DSpark"
def __init__(self, vllm_config: VllmConfig, device: torch.device):
super().__init__(vllm_config, device)
# Whether to sample from the anchor position. When True, uses anchor-as-first
# (N slots, each position predicts the next token). When False, uses 1+N
# fill-in block (anchor is a bonus token).
self.sample_from_anchor = getattr(
self.draft_model_config.hf_config, "sample_from_anchor", True
)
if self.sample_from_anchor:
self.num_query_per_req = self.num_speculative_steps
else:
self.num_query_per_req = 1 + self.num_speculative_steps
# DSpark consumes mean-pooled target aux hidden states at the target
# layers, combined to hidden_size via main_proj. Store that combined
# main_x (hidden_size wide). DSpark does not use the same pre-allocated buffer
# that DeepSeek-V4's MTP uses.
draft_hidden = self.draft_model_config.get_hidden_size()
self.hidden_states = torch.zeros(
self.max_num_tokens, draft_hidden, dtype=self.dtype, device=device
)
self._step_cols = torch.arange(
self.num_speculative_steps, dtype=torch.int32, device=device
)
self._anchor_idx = (
torch.arange(self.max_num_reqs, dtype=torch.int64, device=device)
* self.num_query_per_req
)
# Reduced-vocab probabilistic drafting only; set in load_draft_model.
self._d2t_scatter_index: torch.Tensor | None = None
self._draft_scatter_buf: torch.Tensor | None = None
self._draft_topk: int | None = getattr(
self.draft_model_config.hf_config, "dspark_draft_topk", None
)
def load_draft_model(
self,
target_model: torch.nn.Module,
target_attn_layer_names: set[str],
) -> torch.nn.Module:
model = load_dspark_model(target_model, self.vllm_config)
# Reduced draft vocab: probabilistic rejection sampling indexes draft
# logits by target id, so precompute the draft->target column map and a
# scratch buffer to scatter logits into target vocab before sampling.
if self.draft_logits is not None and model.draft_id_to_target_id is not None:
d2t = model.draft_id_to_target_id
self._d2t_scatter_index = (
torch.arange(d2t.shape[0], device=d2t.device) + d2t
)
# -inf once; the per-step scatter overwrites the draft->target
# columns. Kept separate from draft_logits to avoid aliasing.
self._draft_scatter_buf = torch.full(
(self.max_num_reqs, self.vocab_size),
float("-inf"),
dtype=self.draft_logits.dtype,
device=self.device,
)
return model
def _sample_logits(
self,
logits: torch.Tensor,
idx_map: torch.Tensor,
sample_pos: torch.Tensor,
step: int,
) -> torch.Tensor:
if self.draft_logits is None:
return self.model.map_draft_to_target(logits.argmax(dim=-1))
# Probabilistic sampling and rejection operate in target-vocabulary
# space. A reduced draft vocabulary is scattered into its target rows.
if self._d2t_scatter_index is not None:
assert self._draft_scatter_buf is not None
buf = self._draft_scatter_buf[: logits.shape[0]]
buf.index_copy_(1, self._d2t_scatter_index, logits.to(buf.dtype))
logits = buf
# sample_pos is the predicted token's position Q; the target verifies
# it with the predecessor's Gumbel key (Q-1). Pass Q-1.
return gumbel_sample(
logits,
idx_map,
self.temperature,
self.seeds,
sample_pos - 1,
apply_temperature=True,
logits_cache=self.draft_logits,
logits_cache_col=self._step_cols[step],
use_fp64=self.use_fp64_gumbel,
)
def _sample_sequential(self, num_reqs: int, head_hidden: torch.Tensor) -> None:
if self._draft_topk is not None:
self._sample_sequential_topk(num_reqs, head_hidden)
return
# Sequential Markov sampling over the backbone's output hidden states.
n_spec = self.num_speculative_steps
num_sample = num_reqs * n_spec
# Per-(req, position) head hidden, ordered (req, step).
sample_hidden = head_hidden[self.sample_indices[:num_sample]]
# Draft-vocab logits; sampled ids are remapped to target vocab below.
base_logits = self.model.compute_draft_logits(sample_hidden)
vocab_size = base_logits.shape[-1]
base_logits = base_logits.view(num_reqs, n_spec, vocab_size)
idx_map = self.sample_idx_mapping[:num_sample].view(num_reqs, n_spec)
sample_pos = self.sample_pos[:num_sample].view(num_reqs, n_spec)
# Anchor (bonus) token per request = the input id at query offset 0,
# read via the precomputed persistent index (fixed buffer for capture).
prev = self.input_buffers.input_ids[self._anchor_idx[:num_reqs]]
for i in range(n_spec):
# Sequential stage: Markov bias from the previously sampled token.
markov_embed = self.model.markov_embed(prev)
bias = self.model.markov_bias(markov_embed)
logits_i = base_logits[:, i] + bias
draft_sampled_i = self._sample_logits(
logits_i, idx_map[:, i], sample_pos[:, i], i
)
self.draft_tokens[:num_reqs, i] = draft_sampled_i
prev = draft_sampled_i
def _sample_sequential_topk(self, num_reqs: int, head_hidden: torch.Tensor) -> None:
"""Apply the sequential Markov head only to top-k base-logit candidates.
Candidate selection is done once for all draft positions. At each
sequential step, the selected logits are corrected in place and every
other entry is set to ``-inf``. The normal dense sampling and rejection
paths then consume that truncated distribution unchanged.
"""
assert self._draft_topk is not None
n_spec = self.num_speculative_steps
num_sample = num_reqs * n_spec
sample_hidden = head_hidden[self.sample_indices[:num_sample]]
base_logits = self.model.compute_draft_logits(sample_hidden)
base_logits = base_logits.view(num_reqs, n_spec, -1)
base_values, draft_indices = base_logits.topk(self._draft_topk, dim=-1)
# Reuse the dense backbone output as the normal sampler's input. Fill
# once for all positions, then scatter only the corrected candidates
# during the sequential loop.
base_logits.fill_(float("-inf"))
idx_map = self.sample_idx_mapping[:num_sample].view(num_reqs, n_spec)
sample_pos = self.sample_pos[:num_sample].view(num_reqs, n_spec)
prev = self.input_buffers.input_ids[self._anchor_idx[:num_reqs]]
for i in range(n_spec):
markov_embed = self.model.markov_embed(prev)
logits_i = self.model.apply_markov_bias_gathered(
markov_embed,
base_logits[:, i],
base_values[:, i],
draft_indices[:, i],
)
draft_sampled_i = self._sample_logits(
logits_i, idx_map[:, i], sample_pos[:, i], i
)
self.draft_tokens[:num_reqs, i] = draft_sampled_i
prev = draft_sampled_i
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:
# Full draft step (captured under CUDA graph): parallel backbone forward
# then sequential Markov sampling over its hidden state outputs.
head_hidden = self._run_model(
num_tokens_padded,
attn_metadata,
slot_mappings,
num_tokens_across_dp,
cudagraph_runtime_mode,
)
self._sample_sequential(num_reqs, head_hidden)