vlm-twin-spec-decoding / vllm_patch /mrope_draft_spec.patch
LeoMaxwell's picture
add code, data, results, patches, report
ee3e28a verified
Raw History Blame Contribute Delete
12.9 kB
--- /testessfs10/users/zeyu.zhang/wangyu_ssd/vllm_env/lib/python3.11/site-packages/vllm/v1/spec_decode/llm_base_proposer.py.orig 2026-08-16 17:03:01.543167000 +0800
+++ /testessfs10/users/zeyu.zhang/wangyu_ssd/vllm_env/lib/python3.11/site-packages/vllm/v1/spec_decode/llm_base_proposer.py 2026-08-16 17:03:18.492170389 +0800
@@ -186,6 +186,16 @@
self.mrope_positions = torch.zeros(
(3, self.max_positions + 1), dtype=torch.int64, device=device
)
+ # Draft-model path: slot mapping must be computed from SEQUENCE
+ # indices, which differ from (compressed) M-RoPE position values
+ # whenever the prompt contains an image. self.positions holds the
+ # sequence indices; seq_step tracks them across drafting steps.
+ self.positions = torch.zeros(
+ self.max_positions, dtype=torch.int64, device=device
+ )
+ self.seq_step = torch.zeros(
+ self.max_batch_size, dtype=torch.int64, device=device
+ )
elif self.uses_xdrope_dim > 0 and self.draft_uses_xdrope_dim > 0:
self.xdrope_positions = torch.zeros(
(self.uses_xdrope_dim, self.max_positions + 1),
@@ -334,9 +344,13 @@
)
def _raise_if_mrope(self):
- if self.draft_model_config.uses_mrope:
+ # M-RoPE is supported on the draft-model path: slot mapping is computed
+ # from true sequence indices (self.positions doubles as a seq-index
+ # buffer, see set_inputs_first_pass) while M-RoPE positions are
+ # scattered separately into self.mrope_positions.
+ if self.draft_model_config.uses_mrope and self.parallel_drafting:
raise NotImplementedError(
- "Speculative Decoding with draft models or parallel drafting "
+ "Speculative Decoding with parallel drafting "
"does not support M-RoPE yet"
)
@@ -628,6 +642,11 @@
if self.uses_mrope:
positions = self.mrope_positions[:, token_indices_to_sample]
+ if self.needs_extra_input_slots:
+ # Track true sequence indices alongside M-RoPE positions for
+ # per-step slot mapping (self.positions holds sequence indices
+ # on this path; see set_inputs_first_pass).
+ self.seq_step[:batch_size] = self.positions[token_indices_to_sample]
else:
positions = self.positions[token_indices_to_sample]
hidden_states = hidden_states[token_indices_to_sample]
@@ -776,12 +795,22 @@
) -> torch.Tensor:
"""Update positions, slot mappings, and sequence metadata for the
next draft step. Returns the updated positions tensor."""
- positions_1d = positions[0] if self.uses_mrope else positions
- if self.uses_mrope:
+ mrope_draft_model = self.uses_mrope and self.needs_extra_input_slots
+ if mrope_draft_model:
+ # Slot mapping must advance by SEQUENCE index, not by (compressed)
+ # M-RoPE position value. seq_step is advanced in place by the
+ # fused kernel; M-RoPE positions advance by +1 on all three rows
+ # separately below (drafted tokens are always text-region).
+ positions_1d = self.seq_step[:batch_size]
+ out_pos = self.seq_step[:batch_size]
+ elif self.uses_mrope:
+ positions_1d = positions[0]
out_pos = self.mrope_positions[0, :batch_size]
elif self.uses_xdrope_dim > 0 and self.draft_uses_xdrope_dim > 0:
+ positions_1d = positions
out_pos = self.xdrope_positions[0, :batch_size]
else:
+ positions_1d = positions
out_pos = self.positions[:batch_size]
eagle_step_update_slot_mapping_and_metadata(
positions_1d=positions_1d,
@@ -794,7 +823,14 @@
input_batch_size=input_batch_size,
)
common_attn_metadata.slot_mapping = self._slot_mapping_buffer[:batch_size]
- if self.uses_mrope:
+ if mrope_draft_model:
+ torch.clamp(
+ positions + 1,
+ max=self.max_model_len - 1,
+ out=self.mrope_positions[:, :batch_size],
+ )
+ positions = self.mrope_positions[:, :batch_size]
+ elif self.uses_mrope:
self.mrope_positions[1:, :batch_size] = self.mrope_positions[0, :batch_size]
positions = self.mrope_positions[:, :batch_size]
elif self.uses_xdrope_dim > 0 and self.draft_uses_xdrope_dim > 0:
@@ -898,14 +934,34 @@
if num_rejected_tokens_gpu is not None:
query_end_loc = query_end_loc - num_rejected_tokens_gpu
+ kernel_positions_in = target_positions
+ if self.uses_mrope:
+ # The kernel synthesizes positions as start_pos + j, reading
+ # only target_positions_ptr[query_start_loc[r]] per request.
+ # Slot mapping (step 2 below) consumes self.positions as
+ # SEQUENCE indices, which diverge from compressed M-RoPE
+ # position values on image prompts. Feed the kernel per-request
+ # context lengths so self.positions ends up holding sequence
+ # indices; true M-RoPE positions are scattered into
+ # self.mrope_positions right after the kernel.
+ naive_qlens = cad.query_start_loc[1:] - cad.query_start_loc[:-1]
+ seq_scratch = torch.zeros(
+ total_num_input_tokens, dtype=torch.int64, device=self.device
+ )
+ seq_scratch[cad.query_start_loc[:-1].to(torch.int64)] = (
+ cad.seq_lens - naive_qlens
+ ).to(torch.int64)
+ kernel_positions_in = seq_scratch
+
copy_and_expand_eagle_inputs_kernel[grid](
# (Padded) Inputs from the target model
target_token_ids_ptr=target_token_ids,
- target_positions_ptr=target_positions,
+ target_positions_ptr=kernel_positions_in,
next_token_ids_ptr=next_token_ids, # sampled tokens, one per request
# Outputs to the drafting buffers
out_input_ids_ptr=self.input_ids,
- out_positions_ptr=self.positions, # Doesn't support mrope for now
+ # For M-RoPE this receives sequence indices (see above)
+ out_positions_ptr=self.positions,
out_is_rejected_token_mask_ptr=self.is_rejected_token_mask,
out_is_masked_token_mask_ptr=self.is_masked_token_mask,
out_new_token_indices_ptr=token_indices_to_sample,
@@ -922,6 +978,31 @@
shift_input_ids=self.pass_hidden_states_to_model,
BLOCK_SIZE_TOKENS=BLOCK_SIZE_TOKENS,
)
+
+ self._mm_expand_map = None
+ if self.uses_mrope:
+ # Scatter true M-RoPE positions into the expanded layout.
+ # shift_input_ids is False on this path, so input token i of
+ # request r lands at output slot i + r * extra_slots_per_request.
+ # ORDERING IS LOAD-BEARING: under this formula the rejected-
+ # region inputs collide with the bonus slot, so the bonus-slot
+ # write below must come last to overwrite them. Trailing
+ # rejected output slots keep stale values, which is harmless
+ # (their KV writes are masked and outputs ignored).
+ req_ids = torch.repeat_interleave(
+ self.arange[:batch_size].to(torch.int64), naive_qlens
+ )
+ out_map = (
+ self.arange[:total_num_input_tokens]
+ + req_ids * self.extra_slots_per_request
+ )
+ self.mrope_positions[:, out_map] = target_positions[
+ :, :total_num_input_tokens
+ ]
+ self.mrope_positions[:, token_indices_to_sample.to(torch.int64)] = (
+ target_positions[:, query_end_loc.to(torch.int64)] + 1
+ )
+ self._mm_expand_map = out_map
if self.pass_hidden_states_to_model:
assert self.parallel_drafting_hidden_state_tensor is not None
self.hidden_states[out_hidden_state_mapping] = target_hidden_states
@@ -968,6 +1049,21 @@
if self.supports_mm_inputs:
mm_embeds, is_mm_embed = mm_embed_inputs or (None, None)
+ if (
+ is_mm_embed is not None
+ and getattr(self, "_mm_expand_map", None) is not None
+ ):
+ # Draft-model path: re-align the multimodal mask from the
+ # runner's token layout to the expanded (bonus-slot-inserted)
+ # drafting layout. Bonus/rejected slots are text -> False.
+ expanded = torch.zeros(
+ num_tokens, dtype=is_mm_embed.dtype, device=is_mm_embed.device
+ )
+ expanded[self._mm_expand_map] = is_mm_embed[
+ : self._mm_expand_map.shape[0]
+ ]
+ is_mm_embed = expanded
+
self.inputs_embeds[:num_tokens] = self.model.embed_input_ids(
self.input_ids[:num_tokens],
multimodal_embeddings=mm_embeds,
--- /testessfs10/users/zeyu.zhang/wangyu_ssd/vllm_env/lib/python3.11/site-packages/vllm/v1/worker/gpu_model_runner.py.orig 2026-08-16 17:03:01.546349000 +0800
+++ /testessfs10/users/zeyu.zhang/wangyu_ssd/vllm_env/lib/python3.11/site-packages/vllm/v1/worker/gpu_model_runner.py 2026-08-16 17:03:23.167490000 +0800
@@ -5260,7 +5260,7 @@
if self.supports_mm_inputs and self.drafter.supports_mm_inputs:
mm_embed_inputs = self._gather_mm_embeddings(
scheduler_output,
- shift_computed_tokens=1,
+ shift_computed_tokens=(1 if self.drafter.pass_hidden_states_to_model else 0),
)
else:
mm_embed_inputs = None
--- /testessfs10/users/zeyu.zhang/wangyu_ssd/vllm_env/lib/python3.11/site-packages/vllm/model_executor/layers/quantization/compressed_tensors/utils.py.orig 2026-08-16 17:08:26.513544000 +0800
+++ /testessfs10/users/zeyu.zhang/wangyu_ssd/vllm_env/lib/python3.11/site-packages/vllm/model_executor/layers/quantization/compressed_tensors/utils.py 2026-08-16 17:08:26.531649000 +0800
@@ -55,6 +55,12 @@
if layer_name is None:
return False
+ # Draft models in speculative decoding are built under the "draft_model."
+ # prefix (see DraftModelProposer._get_model), which is unknown to the
+ # checkpoint ignore list -- strip it before matching.
+ if layer_name.startswith("draft_model."):
+ layer_name = layer_name[len("draft_model."):]
+
# layer_name = model.layers.0.self_attn.qkv_proj
# proj_name = qkv_proj
proj_name = layer_name.split(".")[-1]
--- /testessfs10/users/zeyu.zhang/wangyu_ssd/vllm_env/lib/python3.11/site-packages/vllm/v1/sample/rejection_sampler.py.orig 2026-08-16 17:51:54.716036000 +0800
+++ /testessfs10/users/zeyu.zhang/wangyu_ssd/vllm_env/lib/python3.11/site-packages/vllm/v1/sample/rejection_sampler.py 2026-08-16 17:51:54.737130000 +0800
@@ -8,6 +8,10 @@
from typing import TYPE_CHECKING
import torch
+
+import os
+
+_RELAX_LOGTAU = float(os.environ.get("VLLM_SPEC_RELAX_LOGTAU", "0") or "0")
import torch.nn as nn
from vllm.config.model import PROCESSED_LOGPROBS_MODES
@@ -453,6 +457,22 @@
if not sampling_metadata.all_random:
# Rejection sampling for greedy sampling requests.
target_argmax = target_logits.argmax(dim=-1)
+ if _RELAX_LOGTAU > 0.0 and num_tokens > 0 and sampling_metadata.all_greedy:
+ # Relaxed acceptance gate (logit-ratio, CSD-style): accept draft
+ # token d when z(d) >= z(argmax) - logtau (logtau = -ln tau).
+ # Implemented by overwriting target_argmax at gate-passing
+ # positions, so the greedy kernel realizes the relaxed policy;
+ # correction/bonus tokens are unaffected (they only come from
+ # gate-failing positions or the bonus slot).
+ _zmax = target_logits.gather(1, target_argmax.unsqueeze(1)).squeeze(1)
+ _zd = target_logits.gather(
+ 1, draft_token_ids.to(torch.int64).unsqueeze(1)
+ ).squeeze(1)
+ target_argmax = torch.where(
+ _zd >= _zmax - _RELAX_LOGTAU,
+ draft_token_ids.to(target_argmax.dtype),
+ target_argmax,
+ )
rejection_greedy_sample_kernel[(batch_size,)](
output_token_ids,
cu_num_draft_tokens,