[Spec Decode] Add top-k DSpark Markov projection (#49969)

Signed-off-by: Andrii Skliar <[email protected]>
Co-authored-by: Andrii Skliar <[email protected]>
This commit is contained in:
Andrii Skliar
2026-08-04 12:52:53 +00:00
committed by GitHub
co-authored by Andrii Skliar
parent 1eb3694521
commit 5789897aa4
4 changed files with 235 additions and 26 deletions
+58
View File
@@ -0,0 +1,58 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch
from torch import nn
from vllm.model_executor.models.qwen3_dspark import DSparkMarkovHead
def _markov_head(weight: torch.Tensor) -> DSparkMarkovHead:
head = DSparkMarkovHead.__new__(DSparkMarkovHead)
nn.Module.__init__(head)
head.markov_w2 = nn.Linear(
weight.shape[1], weight.shape[0], bias=False, dtype=weight.dtype
)
head.markov_w2.weight.data.copy_(weight)
return head
def test_gathered_markov_bias_overwrites_dense_logits():
weight = torch.arange(21, dtype=torch.float32).view(7, 3) / 10
markov_embed = torch.tensor([[0.5, -1.0, 0.25], [1.0, 0.5, -0.5]])
logits = torch.tensor(
[
[0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7],
[0.7, 0.6, 0.5, 0.4, 0.3, 0.2, 0.1],
]
)
values, index = logits.topk(3, dim=-1)
values = torch.stack((values, torch.zeros_like(values)), dim=1)[:, 0]
expected = values + torch.bmm(weight[index], markov_embed.unsqueeze(-1)).squeeze(-1)
logits.fill_(float("-inf"))
result = _markov_head(weight).apply_bias_gathered(
markov_embed, logits, values, index
)
assert result is logits
torch.testing.assert_close(result.gather(1, index), expected)
selected = torch.zeros_like(result, dtype=torch.bool).scatter_(1, index, True)
assert torch.isneginf(result.masked_select(~selected)).all()
def test_gathered_markov_bias_matches_dense_at_full_vocab():
weight = torch.arange(15, dtype=torch.float32).view(5, 3) / 10
markov_embed = torch.tensor([[0.5, -1.0, 0.25]])
logits = torch.tensor([[0.1, 0.4, -0.2, 0.3, 0.0]])
original = logits.clone()
values, index = logits.topk(logits.shape[-1], dim=-1)
scale = 0.5
logits.fill_(float("-inf"))
result = _markov_head(weight).apply_bias_gathered(
markov_embed, logits, values, index, scale
)
expected = original + markov_embed @ weight.T * scale
torch.testing.assert_close(result, expected)
+59 -1
View File
@@ -289,6 +289,10 @@ class SpeculativeConfig:
during rejection sampling. This comes at the cost of additional GPU memory
usage."""
dspark_draft_topk: int | None = Field(default=None, ge=1)
"""For Qwen3 DSpark drafting, evaluate the Markov projection only for the
top-k base-logit candidates. Requires draft tensor parallel size 1."""
def compute_hash(self) -> str:
"""
WARNING: Whenever a new field is added to this config,
@@ -1032,6 +1036,61 @@ class SpeculativeConfig:
"Inkling MTP currently supports exactly one speculative token"
)
if self.dspark_draft_topk is not None and self.method != "dspark":
raise ValueError("dspark_draft_topk is only supported by DSpark")
dspark_draft_topk = None
if self.method == "dspark":
# DSpark is a semi-autoregressive *block* drafter. A
# speculative length smaller than the checkpoint's block
# feeds the block / Markov-head machinery an unsupported
# layout and yields incorrect (garbled) output rather than
# merely lower acceptance. Require num_speculative_tokens to
# be at least the block size (e.g. 5 or 7 for DeepSeek-V4).
dspark_block_size = getattr(
self.draft_model_config.hf_config,
"dspark_block_size",
None,
)
if (
dspark_block_size is not None
and self.num_speculative_tokens < dspark_block_size
):
raise ValueError(
"DSpark requires num_speculative_tokens >= "
f"dspark_block_size ({dspark_block_size}); got "
f"{self.num_speculative_tokens}. Smaller values "
"produce incorrect output. Use "
f"num_speculative_tokens={dspark_block_size} or "
"larger (e.g. 7)."
)
hf_config = self.draft_model_config.hf_config
dspark_draft_topk = self.dspark_draft_topk
if dspark_draft_topk is None:
dspark_draft_topk = getattr(
hf_config, "dspark_draft_topk", None
)
if dspark_draft_topk is not None:
draft_vocab_size = (
getattr(hf_config, "draft_vocab_size", None)
or hf_config.vocab_size
)
if not 1 <= dspark_draft_topk <= draft_vocab_size:
raise ValueError(
"dspark_draft_topk must be between 1 and the "
f"draft vocabulary size ({draft_vocab_size})"
)
if (
"Qwen3DSparkModel"
not in self.draft_model_config.architectures
):
raise ValueError(
"dspark_draft_topk is only supported by "
"Qwen3DSparkModel"
)
hf_config.dspark_draft_topk = dspark_draft_topk
self.draft_tensor_parallel_size = (
SpeculativeConfig._verify_and_get_draft_tp(
self.target_parallel_config,
@@ -1039,7 +1098,6 @@ class SpeculativeConfig:
self.draft_model_config.hf_config,
)
)
self.draft_model_config.max_model_len = (
SpeculativeConfig._maybe_override_draft_max_model_len(
self.max_model_len,
@@ -77,6 +77,30 @@ class DSparkMarkovHead(nn.Module):
"""Vocab-size transition bias from a Markov embedding ([B, r] -> [B, V])."""
return logits_processor(self.markov_w2, markov_embed)
def apply_bias_gathered(
self,
markov_embed: torch.Tensor,
logits: torch.Tensor,
values: torch.Tensor,
index: torch.Tensor,
scale: float = 1.0,
) -> torch.Tensor:
"""Apply the Markov bias only to selected rows of ``logits``.
The caller initializes ``logits`` to ``-inf`` once for all draft
positions. This method scatters the corrected candidate values into
that dense buffer so the normal sampler sees the truncated proposal.
"""
weight = self.markov_w2.weight[index]
corrected = values.unsqueeze(-1)
corrected.baddbmm_(
weight,
markov_embed.unsqueeze(-1),
beta=1.0,
alpha=scale,
)
return logits.scatter_(1, index, corrected.squeeze(-1))
class Qwen3DSparkModel(DFlashQwen3Model):
"""DFlash Qwen3 backbone + DSpark Markov head."""
@@ -158,6 +182,21 @@ class Qwen3DSparkForCausalLM(DFlashQwen3ForCausalLM):
def markov_bias(self, markov_embed: torch.Tensor) -> torch.Tensor:
return self.model.markov_head.bias(markov_embed, self.logits_processor)
def apply_markov_bias_gathered(
self,
markov_embed: torch.Tensor,
logits: torch.Tensor,
values: torch.Tensor,
index: torch.Tensor,
) -> torch.Tensor:
return self.model.markov_head.apply_bias_gathered(
markov_embed,
logits,
values,
index,
self.logits_processor.scale,
)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
model_weights = {}
includes_embed_tokens = False
@@ -72,6 +72,9 @@ class DSparkSpeculator(DFlashSpeculator):
# 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,
@@ -97,7 +100,43 @@ class DSparkSpeculator(DFlashSpeculator):
)
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,
output_processed_logits=self.draft_logits,
output_processed_logits_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
@@ -120,31 +159,46 @@ class DSparkSpeculator(DFlashSpeculator):
markov_embed = self.model.markov_embed(prev)
bias = self.model.markov_bias(markov_embed)
logits_i = base_logits[:, i] + bias
if self.draft_logits is not None:
# Probabilistic: sample in target vocab (a reduced draft vocab is
# scattered into its target columns; full vocab is already there).
if self._d2t_scatter_index is not None:
assert self._draft_scatter_buf is not None
buf = self._draft_scatter_buf[:num_reqs]
buf.index_copy_(1, self._d2t_scatter_index, logits_i.to(buf.dtype))
logits_i = 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.
draft_sampled_i = gumbel_sample(
logits_i,
idx_map[:, i],
self.temperature,
self.seeds,
sample_pos[:, i] - 1,
apply_temperature=True,
output_processed_logits=self.draft_logits,
output_processed_logits_col=self._step_cols[i],
use_fp64=self.use_fp64_gumbel,
)
else:
draft_sampled_i = self.model.map_draft_to_target(
logits_i.argmax(dim=-1)
)
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