Skip to content
Merged
Show file tree
Hide file tree
Changes from 9 commits
Commits
Show all changes
21 commits
Select commit Hold shift + click to select a range
10cf5f8
[Feature] Universal speculative decoding for heterogeneous vocabulari…
wonderful199082 Mar 26, 2026
6896cbd
fix: raise ValueError when tokenizer lacks unk_token_id in VocabMapping
wan-danfeng Mar 26, 2026
90d145c
fix: remove stray colon in universal_draft condition
wan-danfeng Apr 10, 2026
3bc2b3b
vocab_mapping: fix unk fallback, dynamic space prefix, remove redunda…
wan-danfeng Apr 17, 2026
eb4b2f1
spec_decode: merge UniversalDraftModelProposer into DraftModelProposer
wan-danfeng Apr 29, 2026
9c012e9
Remove redundant functions
wan-danfeng Apr 30, 2026
996ffee
chore: address pre-commit warnings
wan-danfeng May 15, 2026
aa804d7
fix: add use_heterogeneous_vocab flag instead of universal_vocab method
wan-danfeng Jun 6, 2026
53179f5
fix: remove redundant function
wan-danfeng Jun 6, 2026
7b40ced
fix: pre-commit
wan-danfeng Jun 6, 2026
dc05e7f
fix: vocab mapping in probabilistic sampling
wan-danfeng Jun 6, 2026
d8161db
Merge branch 'main' into feat/universal-draft-tli
wan-danfeng Jun 7, 2026
6ab7a00
Merge branch 'main' into feat/universal-draft-tli
wan-danfeng Jun 8, 2026
03fa9d8
Update vllm/v1/worker/gpu_model_runner.py
wan-danfeng Jun 10, 2026
4b82b77
fix: validate greedy draft sampling only when TLI is enabled and add …
wan-danfeng Jun 10, 2026
6b568aa
doc: pre-commit check
wan-danfeng Jun 10, 2026
e595009
Merge branch 'main' into feat/universal-draft-tli
wan-danfeng Jun 10, 2026
3f9a0f1
Merge branch 'main' into feat/universal-draft-tli
benchislett Jun 15, 2026
1323c52
Merge branch 'main' into feat/universal-draft-tli
wan-danfeng Jun 15, 2026
461d8dd
Merge branch 'main' into feat/universal-draft-tli
wan-danfeng Jul 1, 2026
a71ce73
Merge branch 'main' into feat/universal-draft-tli
wan-danfeng Jul 1, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions examples/features/speculative_decoding/spec_decode_offline.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,7 @@ def parse_args():
parser.add_argument("--max-num-seqs", type=int, default=None)
parser.add_argument("--parallel-drafting", action="store_true")
parser.add_argument("--allowed-local-media-path", type=str, default="")
parser.add_argument("--use-heterogeneous-vocab", action="store_true")
return parser.parse_args()


Expand Down Expand Up @@ -135,6 +136,7 @@ def main(args):
"enforce_eager": args.enforce_eager,
"max_model_len": args.max_model_len,
"parallel_drafting": args.parallel_drafting,
"use_heterogeneous_vocab": args.use_heterogeneous_vocab,
}
elif args.method == "mtp":
speculative_config = {
Expand Down
48 changes: 48 additions & 0 deletions tests/v1/spec_decode/test_vocab_mapping.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import pytest
from transformers import AutoTokenizer

from vllm.v1.spec_decode.vocab_mapping import _detect_space_prefix


@pytest.mark.parametrize(
"model_name,expected_prefix",
[
# BPE tokenizer (GPT-2 family) uses Ġ (U+0120)
("HuggingFaceTB/SmolLM2-135M-Instruct", ("Ġ",)),
# SentencePiece tokenizer (LLaMA family) uses ▁ (U+2581)
("TinyLlama/TinyLlama-1.1B-Chat-v1.0", ("▁",)),
# BPE tokenizer (Qwen family) uses Ġ (U+0120)
("Qwen/Qwen2.5-0.5B-Instruct", ("Ġ",)),
],
)
def test_detect_space_prefix_real_tokenizers(model_name, expected_prefix):
tokenizer = AutoTokenizer.from_pretrained(model_name)
result = _detect_space_prefix(tokenizer)
assert result == expected_prefix, (
f"{model_name}: expected {expected_prefix!r}, got {result!r}"
)


def test_detect_space_prefix_fallback_on_failure():
"""When tokenizer lacks encode(), fall back to both known prefixes."""

class BrokenTokenizer:
def encode(self, text, **kwargs):
raise RuntimeError("broken")

result = _detect_space_prefix(BrokenTokenizer())
assert result == ("Ġ", "▁")


def test_detect_space_prefix_empty_encode():
"""When encode returns empty list, fall back."""

class EmptyTokenizer:
def encode(self, text, **kwargs):
return []

result = _detect_space_prefix(EmptyTokenizer())
assert result == ("Ġ", "▁")
20 changes: 18 additions & 2 deletions vllm/config/speculative.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,12 @@ class SpeculativeConfig:
O(2 * tp_size) per token. Only applies to greedy draft selection in
non-tree speculation."""

use_heterogeneous_vocab: bool = False
"""Allow draft and target models to use different vocabularies.
When enabled, builds a token-level intersection at init and constrains
draft logits to shared tokens only (TLI algorithm). Requires
method='draft_model'."""

# Ngram proposer configuration
prompt_lookup_max: int | None = Field(default=None, ge=1)
"""Maximum size of ngram token window when using Ngram proposer, required
Expand Down Expand Up @@ -674,7 +680,11 @@ def __post_init__(self):
self.draft_model_config = ModelConfig(
model=self.model,
runner="draft",
tokenizer=self.target_model_config.tokenizer,
tokenizer=(
self.model
if self.use_heterogeneous_vocab
else self.target_model_config.tokenizer
),
tokenizer_mode=self.target_model_config.tokenizer_mode,
trust_remote_code=self.target_model_config.trust_remote_code,
allowed_local_media_path=self.target_model_config.allowed_local_media_path,
Expand Down Expand Up @@ -1015,7 +1025,13 @@ def _verify_args(self) -> Self:
self.draft_parallel_config
)

self.verify_equal_vocab_size_if_draft_model()
if self.use_heterogeneous_vocab and not self.uses_draft_model():
raise ValueError(
"use_heterogeneous_vocab only works with method='draft_model'"
)

if not self.use_heterogeneous_vocab:
self.verify_equal_vocab_size_if_draft_model()
return self

def verify_equal_vocab_size_if_draft_model(self):
Expand Down
32 changes: 31 additions & 1 deletion vllm/v1/spec_decode/draft_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,10 @@
from vllm.config.utils import replace
from vllm.logger import init_logger
from vllm.model_executor.model_loader import get_model
from vllm.tokenizers.registry import get_tokenizer
from vllm.v1.attention.backends.utils import CommonAttentionMetadata
from vllm.v1.spec_decode.llm_base_proposer import SpecDecodeBaseProposer
from vllm.v1.spec_decode.vocab_mapping import VocabMapping

logger = init_logger(__name__)

Expand All @@ -27,9 +30,36 @@ def __init__(
pass_hidden_states_to_model=False,
runner=runner,
)
self._raise_if_vocab_size_mismatch()
self._raise_if_draft_tp_mismatch()

self.use_heterogeneous_vocab = (
self.speculative_config.use_heterogeneous_vocab
)

spec = self.speculative_config
if self.use_heterogeneous_vocab:
# Heterogeneous vocabularies: build a VocabMapping to translate
# token IDs between the two tokenizers and constrain draft logits
# to the intersection so rejection sampling stays lossless.
target_tokenizer = get_tokenizer(
spec.target_model_config.tokenizer,
trust_remote_code=spec.target_model_config.trust_remote_code,
)
draft_tokenizer = get_tokenizer(
spec.draft_model_config.model,
trust_remote_code=spec.draft_model_config.trust_remote_code,
)
self.vocab_mapping: VocabMapping | None = VocabMapping(
target_tokenizer=target_tokenizer,
draft_tokenizer=draft_tokenizer,
target_vocab_size=spec.target_model_config.get_vocab_size(),
draft_vocab_size=spec.draft_model_config.get_vocab_size(),
device=device,
)
else:
self._raise_if_vocab_size_mismatch()
self.vocab_mapping = None

def _raise_if_vocab_size_mismatch(self):
self.speculative_config.verify_equal_vocab_size_if_draft_model()

Expand Down
21 changes: 21 additions & 0 deletions vllm/v1/spec_decode/llm_base_proposer.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,10 @@
self.speculative_config.use_local_argmax_reduction
)

self.use_heterogeneous_vocab: bool = (
self.speculative_config.use_heterogeneous_vocab
)

self.max_batch_size = vllm_config.scheduler_config.max_num_seqs
self.max_num_tokens = vllm_config.scheduler_config.max_num_batched_tokens
self.token_arange_np = np.arange(self.max_num_tokens, dtype=np.int32)
Expand Down Expand Up @@ -398,11 +402,16 @@
"""Greedy-sample draft tokens from hidden states."""
if self.use_local_argmax_reduction:
return self.model.get_top_tokens(hidden_states)
if self.use_heterogeneous_vocab:
Comment thread
depthfirst-app[bot] marked this conversation as resolved.
logits = self.model.compute_logits(hidden_states)
logits = self.vocab_mapping.constrain_draft_logits(logits)
draft_token_ids = logits.argmax(dim=-1)
return self.vocab_mapping.map_draft_to_target_ids(draft_token_ids)
return self.model.compute_logits(hidden_states).argmax(dim=-1)

def _sample_from_logits(

Check failure on line 412 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]

Check failure on line 412 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]
self,
logits: torch.Tensor,

Check failure on line 414 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]

Check failure on line 414 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]
sampling_metadata: SamplingMetadata,
) -> tuple[torch.Tensor, torch.Tensor | None]:
if not self._enable_probabilistic_draft_probs:
Expand Down Expand Up @@ -573,12 +582,16 @@
# tensor.argmax() returns int64 by default.
input_ids = draft_token_ids_list[-1].int()

if self.use_heterogeneous_vocab:
# Map target token IDs to draft vocab space (TLI algorithm)
input_ids = self.vocab_mapping.map_target_to_draft_ids(input_ids)

if not self.constant_draft_positions:
positions = self._update_positions_dependent_metadata(
positions,
common_attn_metadata,
batch_size,
input_batch_size,

Check failure on line 594 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]

Check failure on line 594 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]
block_size,
)

Expand Down Expand Up @@ -705,11 +718,19 @@
cad: CommonAttentionMetadata,
num_rejected_tokens_gpu: torch.Tensor | None,
) -> tuple[int, torch.Tensor, CommonAttentionMetadata]:
# Map target token IDs to draft vocab space (TLI algorithm)
if self.use_heterogeneous_vocab:
target_token_ids = self.vocab_mapping.map_target_to_draft_ids(
target_token_ids
)
next_token_ids = self.vocab_mapping.map_target_to_draft_ids(
next_token_ids
)
if not self.needs_extra_input_slots:
# Default EAGLE pathway: no reshaping of input tensors needed.

Check failure on line 730 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]

Check failure on line 730 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]
# Simply rotate the input ids and leave the positions unchanged,
# Inserting the next token ids at the last slot in each request.
if token_indices_to_sample is None:

Check failure on line 733 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]

Check failure on line 733 in vllm/v1/spec_decode/llm_base_proposer.py

View workflow job for this annotation

GitHub Actions / pre-commit

"SpecDecodeBaseProposer" has no attribute "vocab_mapping" [attr-defined]
token_indices_to_sample = cad.query_start_loc[1:] - 1

num_tokens = target_token_ids.shape[0]
Expand Down
154 changes: 154 additions & 0 deletions vllm/v1/spec_decode/vocab_mapping.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,154 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch

from vllm.logger import init_logger

logger = init_logger(__name__)


def _detect_space_prefix(tokenizer) -> tuple[str, ...]:
Comment thread
benchislett marked this conversation as resolved.
"""Detect the space-prefix character(s) by tokenizing a literal space.

Different tokenizer families mark word-initial spaces differently:
BPE uses 'Ġ' (U+0120), SentencePiece uses '▁' (U+2581). Probing at
runtime avoids hardcoding assumptions and correctly handles mixed-family
pairs (e.g. BPE draft + SentencePiece target).
"""
try:
space_ids = tokenizer.encode(" a", add_special_tokens=False)
if space_ids:
tok_str = tokenizer.convert_ids_to_tokens(space_ids[0])
if (
isinstance(tok_str, str)
and len(tok_str) > 1
and tok_str.endswith("a")
and tok_str[0] not in (" ", " ")
):
return (tok_str[:-1],)
except Exception:
pass
# Fallback: cover both BPE (Ġ U+0120) and SentencePiece (▁ U+2581)
return ("\u0120", "\u2581")


def _normalize_token(token: str, space_prefixes: tuple[str, ...]) -> str:
for prefix in space_prefixes:
if token.startswith(prefix):
return " " + token[len(prefix) :]
return token


def _get_unk_token_id(tokenizer, role: str) -> int:
"""Return a safe fallback token ID for out-of-intersection tokens.

Preferred: unk_token_id → eos_token_id → ValueError.
Checking with ``is not None`` is required because token ID 0 is a valid
(and common) unk ID on many tokenizers; using ``or 0`` would silently
mishandle those cases.
"""
unk = getattr(tokenizer, "unk_token_id", None)
if unk is not None:
return unk
eos = getattr(tokenizer, "eos_token_id", None)
if eos is not None:
logger.warning(
"VocabMapping: %s has no unk_token_id; "
"falling back to eos_token_id=%d for out-of-intersection tokens",
role,
eos,
)
return eos
raise ValueError(
f"VocabMapping: {role} has neither unk_token_id nor eos_token_id; "
"cannot safely map out-of-intersection tokens"
)


class VocabMapping:
def __init__(
self,
target_tokenizer,
draft_tokenizer,
target_vocab_size,
draft_vocab_size,
device,
):
self.target_vocab_size = target_vocab_size
self.draft_vocab_size = draft_vocab_size
self.device = device
self.target_unk_token_id = _get_unk_token_id(
target_tokenizer, "target tokenizer"
)
self.draft_unk_token_id = _get_unk_token_id(draft_tokenizer, "draft tokenizer")

target_prefixes = _detect_space_prefix(target_tokenizer)
draft_prefixes = _detect_space_prefix(draft_tokenizer)

target_vocab = target_tokenizer.get_vocab()
draft_vocab = draft_tokenizer.get_vocab()

target_normalized = {}
for token, tid in target_vocab.items():
norm = _normalize_token(token, target_prefixes)
if norm not in target_normalized:
target_normalized[norm] = tid

draft_normalized = {}
for token, tid in draft_vocab.items():
norm = _normalize_token(token, draft_prefixes)
if norm not in draft_normalized:
draft_normalized[norm] = tid

common_tokens = set(target_normalized.keys()) & set(draft_normalized.keys())

draft_to_target = torch.full((draft_vocab_size,), -1, dtype=torch.long)
target_to_draft = torch.full((target_vocab_size,), -1, dtype=torch.long)
intersection_mask_draft = torch.zeros(draft_vocab_size, dtype=torch.bool)

for norm_token in common_tokens:
t_id = target_normalized[norm_token]
d_id = draft_normalized[norm_token]
if t_id < target_vocab_size and d_id < draft_vocab_size:
draft_to_target[d_id] = t_id
target_to_draft[t_id] = d_id
intersection_mask_draft[d_id] = True

self.draft_to_target_ids = draft_to_target.to(device)
self.target_to_draft_ids = target_to_draft.to(device)
self.intersection_mask_draft = intersection_mask_draft.to(device)
self.intersection_size = int(intersection_mask_draft.sum().item())

logger.info(
"VocabMapping initialized: target_vocab=%d, draft_vocab=%d, "
"intersection=%d (%.1f%% of draft, %.1f%% of target)",
target_vocab_size,
draft_vocab_size,
self.intersection_size,
100.0 * self.intersection_size / max(draft_vocab_size, 1),
100.0 * self.intersection_size / max(target_vocab_size, 1),
)

if self.intersection_size < 100:
logger.warning(
"Very small vocabulary intersection (%d tokens).",
self.intersection_size,
)

def map_target_to_draft_ids(self, target_ids):
draft_ids = self.target_to_draft_ids[target_ids] # new tensor; no clone needed
missing = draft_ids == -1
if missing.any():
draft_ids[missing] = self.draft_unk_token_id
return draft_ids.to(target_ids.dtype)

def map_draft_to_target_ids(self, draft_ids):
target_ids = self.draft_to_target_ids[draft_ids] # new tensor; no clone needed
missing = target_ids == -1
if missing.any():
target_ids[missing] = self.target_unk_token_id
return target_ids.to(draft_ids.dtype)

def constrain_draft_logits(self, logits):
# masked_fill returns a new tensor; no clone needed
return logits.masked_fill(~self.intersection_mask_draft, float("-inf"))
1 change: 1 addition & 0 deletions vllm/v1/worker/gpu_model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -559,6 +559,7 @@ def __init__(
from vllm.v1.spec_decode.ngram_proposer import NgramProposer

self.drafter = NgramProposer(self.vllm_config)

Comment thread
wan-danfeng marked this conversation as resolved.
Outdated
elif self.speculative_config.uses_draft_model():
self.drafter = DraftModelProposer(
vllm_config=self.vllm_config,
Expand Down
Loading