From 27e0057aeda6bc443069c20fdf2f3cc95ed892f3 Mon Sep 17 00:00:00 2001 From: Luciano Martins Date: Wed, 6 May 2026 11:39:29 -0300 Subject: [PATCH] [Spec Decode] Add Gemma4 MTP speculative decoding support (#41745) Signed-off-by: Luciano Martins Co-authored-by: Luciano Martins --- tests/models/registry.py | 6 + tests/v1/e2e/spec_decode/test_spec_decode.py | 47 +- vllm/config/speculative.py | 20 + vllm/model_executor/models/gemma4_mtp.py | 603 ++++++++++++++++++ vllm/model_executor/models/registry.py | 1 + .../model_arch_config_convertor.py | 13 + vllm/v1/spec_decode/gemma4.py | 335 ++++++++++ vllm/v1/spec_decode/llm_base_proposer.py | 136 ++-- vllm/v1/worker/gpu_model_runner.py | 30 +- 9 files changed, 1120 insertions(+), 71 deletions(-) create mode 100644 vllm/model_executor/models/gemma4_mtp.py create mode 100644 vllm/v1/spec_decode/gemma4.py diff --git a/tests/models/registry.py b/tests/models/registry.py index 1833c85a8cb..e50b0a8de4d 100644 --- a/tests/models/registry.py +++ b/tests/models/registry.py @@ -1494,6 +1494,12 @@ _SPECULATIVE_DECODING_EXAMPLE_MODELS = { trust_remote_code=True, is_available_online=False, ), + "Gemma4MTPModel": _HfExamplesInfo( + "google/gemma-4-E4B-it", + speculative_model="google/gemma-4-E4B-it-assistant", + trust_remote_code=True, + min_transformers_version="5.8.0", + ), "ErnieMTPModel": _HfExamplesInfo( "baidu/ERNIE-4.5-21B-A3B-PT", trust_remote_code=True, diff --git a/tests/v1/e2e/spec_decode/test_spec_decode.py b/tests/v1/e2e/spec_decode/test_spec_decode.py index 4a9c3faa116..2ab2245b790 100644 --- a/tests/v1/e2e/spec_decode/test_spec_decode.py +++ b/tests/v1/e2e/spec_decode/test_spec_decode.py @@ -729,15 +729,20 @@ def test_eagle_correctness_heavy( False, 0.20, ), # hybrid + MTP, ref: ~34%-35% + ( + ("mtp", "google/gemma-4-E4B-it", 1, "google/gemma-4-E4B-it-assistant"), + False, + 0.50, + ), # gemma4 MTP with assistant model, ref: ~62% ], - ids=["mimo", "deepseek", "qwen3_5-hybrid"], + ids=["mimo", "deepseek", "qwen3_5-hybrid", "gemma4-e4b"], ) @single_gpu_only @large_gpu_mark(min_gb=20) def test_mtp_correctness( monkeypatch: pytest.MonkeyPatch, sampling_config: SamplingParams, - model_setup: tuple[str, str, int], + model_setup: tuple[str, str, int] | tuple[str, str, int, str], mm_enabled: bool, expected_accuracy_threshold: float, ): @@ -753,7 +758,11 @@ def test_mtp_correctness( with monkeypatch.context() as m: m.setenv("VLLM_MLA_DISABLE", "1") - method, model_name, tp_size = model_setup + if len(model_setup) == 4: + method, model_name, tp_size, draft_model = model_setup + else: + method, model_name, tp_size = model_setup + draft_model = None _skip_if_insufficient_gpus_for_tp(tp_size) if "Qwen3.5" in model_name and os.environ.get("VLLM_USE_V2_MODEL_RUNNER"): @@ -764,13 +773,22 @@ def test_mtp_correctness( attn_backend = "TRITON_ATTN" if current_platform.is_rocm() else "auto" - # Qwen3.5 is a VLM; without this, profile_run runs the ViT warmup - # and peaks well above the 18GB MIG slice used by one of the CI - # lanes. This test only exercises text generation, so the vision - # tower is never needed. + # Skip multimodal profiling for models that don't need it in this test. extra_kwargs: dict[str, Any] = {} if "Qwen3.5" in model_name: extra_kwargs["limit_mm_per_prompt"] = {"image": 0, "video": 0} + elif "gemma-4" in model_name: + extra_kwargs["limit_mm_per_prompt"] = {"image": 0, "audio": 0} + + if draft_model is not None and "gemma-4" in draft_model: + import transformers + from packaging.version import Version + + if Version(transformers.__version__) < Version("5.8.0"): + pytest.skip( + "Gemma4 MTP assistant requires transformers>=5.8.0, " + f"got {transformers.__version__}" + ) ref_llm = LLM( model=model_name, @@ -788,15 +806,20 @@ def test_mtp_correctness( torch.accelerator.empty_cache() cleanup_dist_env_and_memory() + speculative_config: dict[str, Any] = { + "method": method, + "num_speculative_tokens": 1, + "max_model_len": 2048, + } + if draft_model is not None: + speculative_config["model"] = draft_model + speculative_config["num_speculative_tokens"] = 2 + spec_llm = LLM( model=model_name, trust_remote_code=True, tensor_parallel_size=tp_size, - speculative_config={ - "method": method, - "num_speculative_tokens": 1, - "max_model_len": 2048, - }, + speculative_config=speculative_config, max_model_len=2048, attention_backend=attn_backend, **extra_kwargs, diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index 0d52f396751..e2904de9aa7 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -50,6 +50,7 @@ MTPModelTypes = Literal[ "pangu_ultra_moe_mtp", "step3p5_mtp", "hy_v3_mtp", + "gemma4_mtp", ] NgramGPUTypes = Literal["ngram_gpu"] DFlashModelTypes = Literal["dflash"] @@ -491,6 +492,17 @@ class SpeculativeConfig: {"n_predict": n_predict, "architectures": ["HYV3MTPModel"]} ) + if hf_config.model_type == "gemma4_assistant": + hf_config.model_type = "gemma4_mtp" + text_config = getattr(hf_config, "text_config", hf_config) + # The assistant runs all decoder layers in a single forward + # call to produce one draft token, so n_predict=1. + # num_kv_shared_layers must be 0: cross-model KV sharing is + # set up by the proposer after model construction. + if hasattr(text_config, "num_kv_shared_layers"): + text_config.num_kv_shared_layers = 0 + hf_config.update({"n_predict": 1, "architectures": ["Gemma4MTPModel"]}) + return hf_config def __post_init__(self): @@ -1040,6 +1052,14 @@ class SpeculativeConfig: slots_per_req += 1 return slots_per_req + def use_gemma4_mtp(self) -> bool: + return ( + self.method == "mtp" + and self.draft_model_config is not None + and getattr(self.draft_model_config.hf_config, "model_type", None) + == "gemma4_mtp" + ) + def use_eagle(self) -> bool: return self.method in ("eagle", "eagle3", "mtp", "dflash") diff --git a/vllm/model_executor/models/gemma4_mtp.py b/vllm/model_executor/models/gemma4_mtp.py new file mode 100644 index 00000000000..c294ffc6f9a --- /dev/null +++ b/vllm/model_executor/models/gemma4_mtp.py @@ -0,0 +1,603 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Inference-only Gemma4 MTP (Multi-Token Prediction) model. + +The Gemma4 assistant model is a lightweight decoder that shares KV cache +with the target (backbone) model. All assistant decoder layers are +KV-shared: they only have Q projections (no K/V projections or norms), +and read K/V from the target model's cache at runtime. + +Checkpoint layout (``gemma4_assistant``):: + + model.embed_tokens.* -- token embeddings + model.layers.{i}.* -- decoder layers (Q-only attention + MLP) + model.norm.* -- final RMSNorm + pre_projection.* -- Linear(2 * backbone_hidden_size, hidden_size) + post_projection.* -- Linear(hidden_size, backbone_hidden_size) + lm_head.* -- language model head (tied to embed_tokens) + masked_embedding.centroids.* -- centroid projection (when use_ordered_embeddings) + masked_embedding.token_ordering -- token-to-centroid mapping buffer +""" + +from collections.abc import Iterable + +import torch +from torch import nn + +from vllm.compilation.decorators import support_torch_compile +from vllm.config import CacheConfig, VllmConfig +from vllm.distributed import ( + get_tensor_model_parallel_world_size, + tensor_model_parallel_all_gather, +) +from vllm.logger import init_logger +from vllm.model_executor.layers.attention import Attention +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ( + ColumnParallelLinear, + RowParallelLinear, +) +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.quantization import QuantizationConfig +from vllm.model_executor.layers.rotary_embedding import get_rope +from vllm.model_executor.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, +) +from vllm.model_executor.model_loader.weight_utils import default_weight_loader +from vllm.sequence import IntermediateTensors + +from .gemma4 import Gemma4MLP, _get_text_config +from .utils import ( + AutoWeightsLoader, + WeightsMapper, + extract_layer_index, + maybe_prefix, +) + +logger = init_logger(__name__) + + +class Gemma4MTPMaskedEmbedder(nn.Module): + """Sparse logit computation via centroid-based vocabulary masking. + + Instead of computing logits against the full vocabulary, projects + hidden states to centroid scores, selects top-K centroids, and + computes logits only for the ~top_k * (vocab_size / num_centroids) + tokens belonging to those centroids. + """ + + token_ordering: torch.Tensor + + def __init__( + self, + hidden_size: int, + vocab_size: int, + num_centroids: int, + centroid_intermediate_top_k: int, + ) -> None: + super().__init__() + self.hidden_size = hidden_size + self.vocab_size = vocab_size + self.num_centroids = num_centroids + self.centroid_intermediate_top_k = centroid_intermediate_top_k + self.vocab_size_per_centroid = vocab_size // num_centroids + self.num_selected = centroid_intermediate_top_k * self.vocab_size_per_centroid + + self.centroids = nn.Linear(hidden_size, num_centroids, bias=False) + self.register_buffer( + "token_ordering", + torch.empty(vocab_size, dtype=torch.long), + ) + + def _select_and_score( + self, + hidden_states: torch.Tensor, + lm_head_weight: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Centroid selection + sparse dot product. + + Returns: + logits: (num_tokens, num_selected) sparse logits. + indices: (num_tokens, num_selected) corresponding vocab indices. + """ + num_tokens = hidden_states.shape[0] + _, top_k_indices = torch.topk( + self.centroids(hidden_states), + k=self.centroid_intermediate_top_k, + dim=-1, + ) + clusters = self.token_ordering.view( + self.num_centroids, + self.vocab_size_per_centroid, + ) + selected = clusters[top_k_indices] + embeddings = lm_head_weight[selected.reshape(-1)].view( + num_tokens, + self.num_selected, + self.hidden_size, + ) + logits = torch.einsum("td,tsd->ts", hidden_states, embeddings) + return logits, selected.view(num_tokens, -1) + + def forward( + self, + hidden_states: torch.Tensor, + lm_head_weight: torch.Tensor, + ) -> torch.Tensor: + """Full-vocab logits with non-selected positions masked to -inf.""" + logits, indices = self._select_and_score(hidden_states, lm_head_weight) + output = torch.full( + (hidden_states.shape[0], self.vocab_size), + fill_value=torch.finfo(hidden_states.dtype).min, + dtype=hidden_states.dtype, + device=hidden_states.device, + ) + return output.scatter_(-1, indices, logits) + + def get_top_tokens( + self, + hidden_states: torch.Tensor, + lm_head_weight: torch.Tensor, + ) -> torch.Tensor: + """Sparse argmax — returns vocab token IDs without full-vocab tensor.""" + logits, indices = self._select_and_score(hidden_states, lm_head_weight) + return indices.gather(-1, logits.argmax(-1, keepdim=True)).squeeze(-1) + + +class Gemma4MTPAttention(nn.Module): + """Q-only attention for Gemma4 MTP layers. + + K/V come from the target model's KV cache via + ``kv_sharing_target_layer_name`` (set by the proposer after + model construction). + """ + + def __init__( + self, + config, + hidden_size: int, + num_heads: int, + num_kv_heads: int, + head_dim: int, + max_position_embeddings: int, + cache_config: CacheConfig | None = None, + quant_config: QuantizationConfig | None = None, + attn_logits_soft_cap: float | None = None, + prefix: str = "", + ) -> None: + super().__init__() + self.config = config + self.hidden_size = hidden_size + + tp_size = get_tensor_model_parallel_world_size() + self.total_num_heads = num_heads + self.num_heads = self.total_num_heads // tp_size + self.total_num_kv_heads = num_kv_heads + self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size) + self.head_dim = head_dim + self.q_size = self.num_heads * self.head_dim + self.scaling = 1.0 + + self.q_proj = ColumnParallelLinear( + hidden_size, + self.total_num_heads * self.head_dim, + bias=config.attention_bias, + quant_config=None, + prefix=f"{prefix}.q_proj", + ) + self.o_proj = RowParallelLinear( + self.total_num_heads * self.head_dim, + hidden_size, + bias=config.attention_bias, + quant_config=None, + prefix=f"{prefix}.o_proj", + ) + self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps) + + layer_idx = extract_layer_index(prefix) + layer_type = config.layer_types[layer_idx] + self.is_sliding = layer_type == "sliding_attention" + sliding_window = config.sliding_window if self.is_sliding else None + + if layer_type in config.rope_parameters: + rope_parameters = dict(config.rope_parameters[layer_type]) + else: + rope_parameters = dict(config.rope_parameters.copy()) + if self.is_sliding: + rope_parameters["rope_theta"] = getattr( + config, "rope_local_base_freq", 10000.0 + ) + + self.rotary_emb = get_rope( + self.head_dim, + max_position=max_position_embeddings, + rope_parameters=rope_parameters, + is_neox_style=True, + ) + + # kv_sharing_target_layer_name is set after model construction + # by Gemma4Proposer._setup_gemma4_kv_sharing(). + self.is_kv_shared_layer = True + self.attn = Attention( + self.num_heads, + self.head_dim, + self.scaling, + num_kv_heads=self.num_kv_heads, + cache_config=cache_config, + quant_config=quant_config, + logits_soft_cap=attn_logits_soft_cap, + per_layer_sliding_window=sliding_window, + prefix=f"{prefix}.attn", + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + q, _ = self.q_proj(hidden_states) + + q = q.unflatten(-1, (self.num_heads, self.head_dim)) + q = self.q_norm(q) + q = q.flatten(-2, -1) + + q, _ = self.rotary_emb(positions, q, None) + + # Attention reads K/V from the target's cache via KV sharing; + # these dummy tensors are never consumed but required by the API. + num_tokens = q.shape[0] + kv_dummy = torch.empty( + num_tokens, + self.num_kv_heads * self.head_dim, + dtype=q.dtype, + device=q.device, + ) + attn_output = self.attn(q, kv_dummy, kv_dummy) + output, _ = self.o_proj(attn_output) + return output + + +class Gemma4MTPDecoderLayer(nn.Module): + def __init__( + self, + config, + cache_config: CacheConfig | None = None, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + self.hidden_size = config.hidden_size + + layer_idx = extract_layer_index(prefix) + layer_type = config.layer_types[layer_idx] + is_full_attention = layer_type == "full_attention" + head_dim = ( + getattr(config, "global_head_dim", config.head_dim) + if is_full_attention + else config.head_dim + ) + + self.self_attn = Gemma4MTPAttention( + config=config, + hidden_size=self.hidden_size, + num_heads=config.num_attention_heads, + num_kv_heads=config.num_key_value_heads, + head_dim=head_dim, + max_position_embeddings=config.max_position_embeddings, + cache_config=cache_config, + quant_config=quant_config, + attn_logits_soft_cap=getattr(config, "attn_logit_softcapping", None), + prefix=f"{prefix}.self_attn", + ) + + text_config = _get_text_config(config) + self.mlp = Gemma4MLP( + hidden_size=self.hidden_size, + intermediate_size=text_config.intermediate_size, + hidden_activation=text_config.hidden_activation, + quant_config=None, + prefix=f"{prefix}.mlp", + ) + + self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.pre_feedforward_layernorm = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.post_feedforward_layernorm = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + + self.register_buffer("layer_scalar", torch.ones(1)) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + residual: torch.Tensor | None, + **kwargs, + ) -> tuple[torch.Tensor, torch.Tensor]: + residual = hidden_states + hidden_states = self.input_layernorm(residual) + + hidden_states = self.self_attn( + positions=positions, + hidden_states=hidden_states, + **kwargs, + ) + + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = hidden_states + residual + residual = hidden_states + + hidden_states = self.pre_feedforward_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + + hidden_states = self.post_feedforward_layernorm(hidden_states) + hidden_states = hidden_states + residual + + hidden_states = hidden_states * self.layer_scalar + return hidden_states, None + + +class Gemma4MultiTokenPredictor(nn.Module): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super().__init__() + + config = vllm_config.speculative_config.draft_model_config.hf_config + text_config = _get_text_config(config) + self.config = text_config + + self.hidden_size = text_config.hidden_size + self.backbone_hidden_size = getattr( + config, "backbone_hidden_size", self.hidden_size + ) + self.vocab_size = text_config.vocab_size + self.num_mtp_layers = text_config.num_hidden_layers + + self.embed_tokens = VocabParallelEmbedding( + self.vocab_size, + self.hidden_size, + ) + + self.pre_projection = ColumnParallelLinear( + 2 * self.backbone_hidden_size, + self.hidden_size, + bias=False, + gather_output=True, + prefix=f"{prefix}.pre_projection", + ) + + self.post_projection = RowParallelLinear( + self.hidden_size, + self.backbone_hidden_size, + bias=False, + input_is_parallel=False, + prefix=f"{prefix}.post_projection", + ) + + self.layers = nn.ModuleList( + Gemma4MTPDecoderLayer( + text_config, + cache_config=vllm_config.cache_config, + quant_config=vllm_config.quant_config, + prefix=f"{prefix}.layers.{idx}", + ) + for idx in range(self.num_mtp_layers) + ) + + self.norm = RMSNorm(self.hidden_size, eps=text_config.rms_norm_eps) + + # After embedding sharing, embed_tokens is replaced with the + # target model's backbone-dim embedding. Scale by + # sqrt(backbone_hidden_size) to match the target's convention. + self.register_buffer( + "normalizer", + torch.tensor(self.backbone_hidden_size**0.5), + persistent=False, + ) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embed_tokens(input_ids) * self.normalizer + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + stacked_params_mapping = [ + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + + params_dict = dict(self.named_parameters()) + params_dict.update(dict(self.named_buffers())) + loaded_params: set[str] = set() + + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + name = name.replace(weight_name, param_name) + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict: + continue + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight, shard_id) + break + else: + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict: + continue + param = params_dict[name] + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight) + loaded_params.add(name) + + return loaded_params + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + hidden_states: torch.Tensor, + intermediate_tensors: IntermediateTensors | None = None, + inputs_embeds: torch.Tensor | None = None, + spec_step_idx: int = 0, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Returns (draft_hidden_states, backbone_hidden_states). + + draft_hidden_states: draft-dim, used by compute_logits via lm_head. + backbone_hidden_states: backbone-dim, stored in the proposer's + hidden-state buffer and fed back as input to the next step. + """ + if inputs_embeds is None: + inputs_embeds = self.embed_input_ids(input_ids) + + combined = torch.cat([inputs_embeds, hidden_states], dim=-1) + hidden_states, _ = self.pre_projection(combined) + + residual = None + for layer in self.layers: + hidden_states, residual = layer( + positions=positions, + hidden_states=hidden_states, + residual=residual, + ) + + draft_hidden_states = self.norm(hidden_states) + + backbone_hidden_states, _ = self.post_projection(draft_hidden_states) + return draft_hidden_states, backbone_hidden_states + + +@support_torch_compile +class Gemma4MTP(nn.Module): + """Gemma4 Multi-Token Prediction model for speculative decoding. + + forward() returns (draft_hidden_states, backbone_hidden_states). + The proposer uses draft_hidden_states for compute_logits (via + the draft-dim lm_head) and backbone_hidden_states for the + hidden-state feedback buffer. + """ + + has_own_lm_head = True + + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_prefix={ + "pre_projection.": "model.pre_projection.", + "post_projection.": "model.post_projection.", + }, + ) + + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super().__init__() + config = vllm_config.speculative_config.draft_model_config.hf_config + text_config = _get_text_config(config) + self.config = config + + self.model = Gemma4MultiTokenPredictor( + vllm_config=vllm_config, + prefix=maybe_prefix(prefix, "draft_model"), + ) + + # lm_head operates in draft-dim. Tied to embed_tokens at init + # so load_weights populates both from a single checkpoint entry. + # After embedding sharing, lm_head.weight still references the + # original draft-dim tensor. + self.lm_head = ParallelLMHead( + text_config.vocab_size, + text_config.hidden_size, + prefix=maybe_prefix(prefix, "lm_head"), + ) + if getattr(config, "tie_word_embeddings", True): + self.lm_head.weight = self.model.embed_tokens.weight + + self.logits_processor = LogitsProcessor( + text_config.vocab_size, + soft_cap=getattr(text_config, "final_logit_softcapping", None), + ) + + if getattr(config, "use_ordered_embeddings", False): + num_centroids = getattr(config, "num_centroids", 2048) + top_k = getattr(config, "centroid_intermediate_top_k", 32) + self.masked_embedding = Gemma4MTPMaskedEmbedder( + hidden_size=text_config.hidden_size, + vocab_size=text_config.vocab_size, + num_centroids=num_centroids, + centroid_intermediate_top_k=top_k, + ) + logger.info( + "Gemma4 MTP: centroids masking enabled " + "(num_centroids=%d, top_k=%d, active_tokens=%d/%d).", + num_centroids, + top_k, + top_k * (text_config.vocab_size // num_centroids), + text_config.vocab_size, + ) + else: + self.masked_embedding = None + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.model.embed_input_ids(input_ids) + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + hidden_states: torch.Tensor, + intermediate_tensors: IntermediateTensors | None = None, + inputs_embeds: torch.Tensor | None = None, + spec_step_idx: int = 0, + **kwargs: object, + ) -> tuple[torch.Tensor, torch.Tensor]: + return self.model( + input_ids, + positions, + hidden_states, + intermediate_tensors, + inputs_embeds, + spec_step_idx, + ) + + def _get_full_lm_head_weight(self) -> torch.Tensor: + lm_head_weight = self.lm_head.weight + tp_size = get_tensor_model_parallel_world_size() + if tp_size > 1: + lm_head_weight = tensor_model_parallel_all_gather( + lm_head_weight, + dim=0, + ) + return lm_head_weight[: self.masked_embedding.vocab_size] + + def compute_logits( + self, + hidden_states: torch.Tensor, + spec_step_idx: int = 0, + ) -> torch.Tensor | None: + if self.masked_embedding is not None: + return self.masked_embedding( + hidden_states, + self._get_full_lm_head_weight(), + ) + return self.logits_processor(self.lm_head, hidden_states) + + def get_top_tokens( + self, + hidden_states: torch.Tensor, + ) -> torch.Tensor: + """Sparse argmax via centroids masking. Returns token IDs directly.""" + return self.masked_embedding.get_top_tokens( + hidden_states, + self._get_full_lm_head_weight(), + ) + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) diff --git a/vllm/model_executor/models/registry.py b/vllm/model_executor/models/registry.py index f82f7d00555..d38cd63b90b 100644 --- a/vllm/model_executor/models/registry.py +++ b/vllm/model_executor/models/registry.py @@ -601,6 +601,7 @@ _SPECULATIVE_DECODING_MODELS = { "EagleDeepSeekMTPModel": ("deepseek_eagle", "EagleDeepseekV3ForCausalLM"), "DeepSeekMTPModel": ("deepseek_mtp", "DeepSeekMTP"), "DeepSeekV4MTPModel": ("deepseek_v4_mtp", "DeepSeekV4MTP"), + "Gemma4MTPModel": ("gemma4_mtp", "Gemma4MTP"), "ErnieMTPModel": ("ernie_mtp", "ErnieMTP"), "ExaoneMoeMTP": ("exaone_moe_mtp", "ExaoneMoeMTP"), "Exaone4_5_MTP": ("exaone4_5_mtp", "Exaone4_5_MTP"), diff --git a/vllm/transformers_utils/model_arch_config_convertor.py b/vllm/transformers_utils/model_arch_config_convertor.py index 79f28d815e2..35fa1313d1e 100644 --- a/vllm/transformers_utils/model_arch_config_convertor.py +++ b/vllm/transformers_utils/model_arch_config_convertor.py @@ -512,6 +512,18 @@ class LongCatFlashMTPModelArchConfigConvertor(ModelArchConfigConvertorBase): return getattr(self.hf_text_config, "num_nextn_predict_layers", 1) +class Gemma4MTPModelArchConfigConvertor(ModelArchConfigConvertorBase): + def get_hidden_size(self) -> int: + # The speculator buffer must match the backbone (target) model's + # hidden dimension, not the draft model's smaller dimension. + return getattr( + self.hf_config, "backbone_hidden_size", super().get_hidden_size() + ) + + def get_num_hidden_layers(self) -> int: + return getattr(self.hf_text_config, "num_hidden_layers", 0) + + class Gemma4ModelArchConfigConvertor(ModelArchConfigConvertorBase): def is_mm_prefix_lm(self) -> bool: return ( @@ -541,6 +553,7 @@ MODEL_ARCH_CONFIG_CONVERTORS = { "falcon": FalconModelArchConfigConvertor, "gemma4": Gemma4ModelArchConfigConvertor, "gemma4_text": Gemma4ModelArchConfigConvertor, + "gemma4_mtp": Gemma4MTPModelArchConfigConvertor, "RefinedWeb": FalconModelArchConfigConvertor, "RefinedWebModel": FalconModelArchConfigConvertor, "nemotron-nas": NemotronNasModelArchConfigConvertor, diff --git a/vllm/v1/spec_decode/gemma4.py b/vllm/v1/spec_decode/gemma4.py new file mode 100644 index 00000000000..b0a02774faf --- /dev/null +++ b/vllm/v1/spec_decode/gemma4.py @@ -0,0 +1,335 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Gemma4 MTP (Multi-Token Prediction) proposer for speculative decoding. + +The Gemma4 assistant model runs all decoder layers per draft step +(producing one token), and all its attention layers share KV cache +with the target model via cross-model KV sharing. +""" + +from collections import defaultdict +from copy import copy + +import torch +import torch.nn as nn + +from vllm.config import VllmConfig, get_layers_from_vllm_config, replace +from vllm.logger import init_logger +from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase +from vllm.v1.attention.backend import CommonAttentionMetadata +from vllm.v1.kv_cache_interface import ( + KVCacheConfig, + KVCacheSpec, + UniformTypeKVCacheSpecs, +) +from vllm.v1.spec_decode.llm_base_proposer import SpecDecodeBaseProposer +from vllm.v1.worker.utils import AttentionGroup + +logger = init_logger(__name__) + + +class Gemma4Proposer(SpecDecodeBaseProposer): + def __init__( + self, + vllm_config: VllmConfig, + device: torch.device, + runner=None, + ): + super().__init__( + vllm_config, + device, + pass_hidden_states_to_model=True, + runner=runner, + ) + # All draft steps predict from the same position (the last + # target-model position), so positions and seq_lens must not + # advance between steps. + self.constant_draft_positions = True + + # Per-group block tables for multi-group KV cache models. + # Populated by gpu_model_runner during _prepare_inputs. + self._per_group_block_tables: dict[int, torch.Tensor] = {} + + # Centroids CUDA graphs — populated in load_model if centroids + # masking is active. _centroids_sizes is pre-sorted for fast + # lookup in _greedy_sample. + self._centroids_sizes: list[int] = [] + self._centroids_graphs: dict[int, torch.cuda.CUDAGraph] = {} + self._centroids_inputs: dict[int, torch.Tensor] = {} + self._centroids_outputs: dict[int, torch.Tensor] = {} + + def set_per_group_block_table(self, gid: int, block_table: torch.Tensor) -> None: + self._per_group_block_tables[gid] = block_table + + def model_returns_tuple(self) -> bool: + # forward() returns (draft_hidden_states, backbone_hidden_states). + # The proposer uses draft_hidden_states for compute_logits and + # backbone_hidden_states for the hidden-state feedback buffer. + return True + + def build_per_group_and_layer_attn_metadata( + self, + common_attn_metadata: CommonAttentionMetadata, + draft_index: int = 0, + ) -> tuple[list[object], dict[str, object]]: + """Build attention metadata using the correct block table per group. + + Gemma4 has multiple KV cache groups (sliding vs full attention) + with different block tables. The base class receives a single + common_attn_metadata whose block_table belongs to one group. + We swap in the correct block table for each draft attention group. + """ + per_group_attn_metadata: list[object] = [] + per_layer_attn_metadata: dict[str, object] = {} + for attn_group in self.draft_attn_groups: + gid = attn_group.kv_cache_group_id + if gid in self._per_group_block_tables: + cm = copy(common_attn_metadata) + cm.block_table_tensor = self._per_group_block_tables[gid] + else: + cm = common_attn_metadata + attn_metadata = attn_group.get_metadata_builder().build_for_drafting( + common_attn_metadata=cm, draft_index=draft_index + ) + per_group_attn_metadata.append(attn_metadata) + for layer_name in attn_group.layer_names: + per_layer_attn_metadata[layer_name] = attn_metadata + return per_group_attn_metadata, per_layer_attn_metadata + + def _greedy_sample(self, hidden_states: torch.Tensor) -> torch.Tensor: + if self._centroids_sizes: + T = hidden_states.shape[0] + for size in self._centroids_sizes: + if size >= T: + self._centroids_inputs[size][:T].copy_(hidden_states) + self._centroids_graphs[size].replay() + return self._centroids_outputs[size][:T].clone() + return self.model.get_top_tokens(hidden_states) + return super()._greedy_sample(hidden_states) + + def _setup_centroids_cuda_graphs(self) -> None: + """Capture CUDA graphs for centroids get_top_tokens at key sizes.""" + masked_emb = self.model.masked_embedding + lm_head_weight = self.model._get_full_lm_head_weight() + + for size in [1, 2, 4, 8, 16, 32, 64]: + static_input = torch.zeros( + size, + masked_emb.hidden_size, + dtype=self.dtype, + device=self.device, + ) + for _ in range(3): + masked_emb.get_top_tokens(static_input, lm_head_weight) + torch.accelerator.synchronize() + + g = torch.cuda.CUDAGraph() + with torch.cuda.graph(g): + static_output = masked_emb.get_top_tokens( + static_input, + lm_head_weight, + ) + self._centroids_graphs[size] = g + self._centroids_inputs[size] = static_input + self._centroids_outputs[size] = static_output + + self._centroids_sizes = sorted(self._centroids_graphs) + logger.info( + "Gemma4 MTP: captured centroids CUDA graphs for sizes %s.", + self._centroids_sizes, + ) + + def _create_draft_vllm_config(self) -> VllmConfig: + """Preserve the target's forced TRITON_ATTN backend for draft layers. + + Gemma4 forces TRITON_ATTN due to heterogeneous head dimensions + (head_dim=256 sliding, global_head_dim=512 full). The base class + resets attention_config.backend to None for draft models, causing + sliding layers to fall back to FLASH_ATTN which cannot handle + KV-shared cache. Override to carry the target's backend through. + """ + base = super()._create_draft_vllm_config() + target_backend = self.vllm_config.attention_config.backend + if target_backend is not None: + base = replace( + base, + attention_config=replace( + base.attention_config, + backend=target_backend, + ), + ) + return base + + def _maybe_share_lm_head(self, target_language_model: nn.Module) -> None: + """Gemma4 MTP always keeps its own draft-dim lm_head. + + The draft model's lm_head operates in draft hidden_size (e.g. 256), + which differs from the target's backbone hidden_size (e.g. 1536). + Sharing would break compute_logits (and centroids masking when + use_ordered_embeddings is enabled). + """ + logger.info( + "Gemma4 MTP: keeping draft model's own lm_head (draft_dim != backbone_dim)." + ) + + def load_model(self, target_model: nn.Module) -> None: + target_attn_layer_names = set( + get_layers_from_vllm_config( + self.vllm_config, + AttentionLayerBase, # type: ignore[type-abstract] + ).keys() + ) + + super().load_model(target_model) + + self._setup_gemma4_kv_sharing(target_attn_layer_names) + + if getattr(self.model, "masked_embedding", None) is not None: + self._setup_centroids_cuda_graphs() + + def validate_same_kv_cache_group(self, kv_cache_config: KVCacheConfig) -> None: + """Draft layers span multiple KV cache groups (sliding + full + attention with different head dimensions), so skip the base + class single-group assertion.""" + + def initialize_attn_backend( + self, + kv_cache_config: KVCacheConfig, + kernel_block_sizes: list[int] | None = None, + ) -> None: + """Create separate AttentionGroup objects per KV cache spec + so that each head-dim variant gets its own metadata builder.""" + all_attn_layers = get_layers_from_vllm_config( + self.vllm_config, + AttentionLayerBase, # type: ignore[type-abstract] + ) + + layer_to_gid: dict[str, int] = {} + layer_to_spec: dict[str, KVCacheSpec] = {} + for gid, group in enumerate(kv_cache_config.kv_cache_groups): + group_spec = group.kv_cache_spec + for ln in group.layer_names: + layer_to_gid[ln] = gid + if isinstance(group_spec, UniformTypeKVCacheSpecs): + if ln in group_spec.kv_cache_specs: + layer_to_spec[ln] = group_spec.kv_cache_specs[ln] + else: + tgt = getattr( + all_attn_layers.get(ln), + "kv_sharing_target_layer_name", + None, + ) + if tgt and tgt in group_spec.kv_cache_specs: + layer_to_spec[ln] = group_spec.kv_cache_specs[tgt] + else: + layer_to_spec[ln] = group_spec + else: + layer_to_spec[ln] = group_spec + + attention_groups: dict[tuple[tuple[str, str], KVCacheSpec], AttentionGroup] = {} + for layer_name in self._draft_attn_layer_names: + if layer_name not in layer_to_spec: + continue + attn_layer = all_attn_layers[layer_name] + attn_backend = attn_layer.get_attn_backend() + spec = layer_to_spec[layer_name] + gid = layer_to_gid[layer_name] + group_key = (attn_backend.full_cls_name(), spec) + + if group_key not in attention_groups: + kernel_block_size = ( + kernel_block_sizes[gid] + if kernel_block_sizes is not None and gid < len(kernel_block_sizes) + else None + ) + attn_group = AttentionGroup( + backend=attn_backend, + layer_names=[layer_name], + kv_cache_spec=spec, + kv_cache_group_id=gid, + ) + attn_group.create_metadata_builders( + self.vllm_config, + self.device, + kernel_block_size=kernel_block_size, + ) + attention_groups[group_key] = attn_group + else: + attention_groups[group_key].layer_names.append(layer_name) + + self.draft_attn_groups = list(attention_groups.values()) + if self.draft_attn_groups: + self.kv_cache_gid = self.draft_attn_groups[0].kv_cache_group_id + self.block_size = ( + self.draft_attn_groups[0] + .get_metadata_builder() + .kv_cache_spec.block_size + ) + else: + self.kv_cache_gid = 0 + self.block_size = kv_cache_config.kv_cache_groups[ + 0 + ].kv_cache_spec.block_size + logger.debug("Using block size %d for drafting layers", self.block_size) + + def _setup_gemma4_kv_sharing( + self, + target_attn_layer_names: set[str], + ) -> None: + """Wire draft layers to share KV with the target model. + + Each draft decoder layer is mapped to the last non-KV-shared + target layer of the same attention type (sliding or full). + """ + draft_config = self.speculative_config.draft_model_config.hf_config + draft_text_config = draft_config.get_text_config() + target_config = self.vllm_config.model_config.hf_config + target_text_config = target_config.get_text_config() + target_layer_types = getattr(target_text_config, "layer_types", []) + + if not (hasattr(self.model, "model") and hasattr(self.model.model, "layers")): + return + + target_num_kv_shared = getattr(target_text_config, "num_kv_shared_layers", 0) + num_non_shared = len(target_layer_types) - target_num_kv_shared + type_to_target_indices: dict[str, list[int]] = defaultdict(list) + for idx, lt in enumerate(target_layer_types[:num_non_shared]): + type_to_target_indices[lt].append(idx) + + target_prefix = "model.layers" + for name in target_attn_layer_names: + if ".layers." in name: + target_prefix = name.split(".layers.")[0] + ".layers" + break + + draft_layer_types = getattr(draft_text_config, "layer_types", []) + for draft_idx, layer in enumerate(self.model.model.layers): + if not hasattr(layer, "self_attn"): + continue + attn = getattr(layer.self_attn, "attn", None) + if attn is None: + continue + + draft_layer_type = ( + draft_layer_types[draft_idx] + if draft_idx < len(draft_layer_types) + else "full_attention" + ) + candidates = type_to_target_indices.get(draft_layer_type, []) + if not candidates: + logger.warning( + "No target layer of type '%s' for draft layer %d", + draft_layer_type, + draft_idx, + ) + continue + + target_idx = candidates[-1] + target_layer_name = f"{target_prefix}.{target_idx}.self_attn.attn" + attn.kv_sharing_target_layer_name = target_layer_name + logger.info( + "Gemma4 MTP: draft layer %d (%s) -> %s", + draft_idx, + draft_layer_type, + target_layer_name, + ) diff --git a/vllm/v1/spec_decode/llm_base_proposer.py b/vllm/v1/spec_decode/llm_base_proposer.py index 7a0775b0a25..8ee349a1cc0 100644 --- a/vllm/v1/spec_decode/llm_base_proposer.py +++ b/vllm/v1/spec_decode/llm_base_proposer.py @@ -105,6 +105,12 @@ class SpecDecodeBaseProposer: ) self.needs_extra_input_slots = self.net_num_new_slots_per_request > 0 + # When True, all draft steps reuse the same position as the + # first step instead of advancing by one each iteration. + # Used by draft models with Q-only attention that share KV + # with the target and always predict from the same position. + self.constant_draft_positions: bool = False + self.parallel_drafting_token_id: int = 0 self.parallel_drafting_hidden_state_tensor: torch.Tensor | None = None if self.parallel_drafting: @@ -389,9 +395,9 @@ class SpecDecodeBaseProposer: return {name: view for name in self._draft_attn_layer_names} def initialize_cudagraph_keys(self, cudagraph_mode: CUDAGraphMode) -> None: - """Initialize cudagraph dispatcher keys for eagle. + """Initialize cudagraph dispatcher keys for the drafter. - Eagle only supports PIECEWISE cudagraphs (via mixed_mode). + Only supports PIECEWISE cudagraphs (via mixed_mode). This should be called after adjust_cudagraph_sizes_for_spec_decode. """ if ( @@ -500,6 +506,12 @@ class SpecDecodeBaseProposer: positions = self.positions[token_indices_to_sample] hidden_states = hidden_states[token_indices_to_sample] + if self.constant_draft_positions: + # Write the sampling positions into the front of the + # positions buffer so that subsequent loop iterations + # (which read via _get_positions) use the correct values. + self.positions[:batch_size] = positions + if any(isinstance(md, TreeAttentionMetadata) for md in per_group_attn_metadata): # Draft using tree attention - requires full logits for top-k logits = self.model.compute_logits(sample_hidden_states) @@ -557,59 +569,25 @@ class SpecDecodeBaseProposer: # cast to int32 is crucial when eagle model is compiled. # tensor.argmax() returns int64 by default. input_ids = draft_token_ids_list[-1].int() - # Use fused kernel for slot mapping and metadata updates. - # Write clamped positions directly into the positions buffer to - # avoid an extra D2D copy for the common (non-mrope) case. - positions_1d = positions[0] if self.uses_mrope else positions - if self.uses_mrope: - out_pos = self.mrope_positions[0, :batch_size] - elif self.uses_xdrope_dim > 0 and self.draft_uses_xdrope_dim > 0: - out_pos = self.xdrope_positions[0, :batch_size] - else: - out_pos = self.positions[:batch_size] - eagle_step_update_slot_mapping_and_metadata( - positions_1d=positions_1d, - block_table_tensor=common_attn_metadata.block_table_tensor, - seq_lens=common_attn_metadata.seq_lens, - block_size=block_size, - max_model_len=self.max_model_len, - out_clamped_positions=out_pos, - out_slot_mapping=self._slot_mapping_buffer[:input_batch_size], - input_batch_size=input_batch_size, - ) - common_attn_metadata.slot_mapping = self._slot_mapping_buffer[:batch_size] - if 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: - self.xdrope_positions[1:, :batch_size] = self.xdrope_positions[ - 0, :batch_size - ] - positions = self.xdrope_positions[0, :batch_size] - else: - positions = self.positions[:batch_size] - # Increment the maximum sequence length. We increment max_seq_len - # unconditionally even though some seq_lens may have been capped above, - # as max_seq_len serves as an upper bound for sequence lengths. - common_attn_metadata.max_seq_len = min( - common_attn_metadata.max_seq_len + 1, self.max_model_len - ) - # Also update the CPU-side shadow; NOTE: this is hacky and should be - # removed in when common_attn_metadata.seq_lens_cpu is deprecated. - if common_attn_metadata._seq_lens_cpu is not None: - common_attn_metadata._seq_lens_cpu += 1 - if common_attn_metadata._num_computed_tokens_cpu is not None: - common_attn_metadata._num_computed_tokens_cpu += 1 - if common_attn_metadata.seq_lens_cpu_upper_bound is not None: - common_attn_metadata.seq_lens_cpu_upper_bound += 1 + if not self.constant_draft_positions: + positions = self._update_positions_dependent_metadata( + positions, + common_attn_metadata, + batch_size, + input_batch_size, + block_size, + ) - # Rebuild attention metadata - _, per_layer_attn_metadata = self.build_per_group_and_layer_attn_metadata( - common_attn_metadata, draft_index=token_index + 1 - ) + # Rebuild attention metadata. When draft positions are constant + # (e.g. Gemma4 MTP), common_attn_metadata is invariant across + # loop iterations so we build once and reuse. + if not self.constant_draft_positions or token_index == 0: + _, per_layer_attn_metadata = ( + self.build_per_group_and_layer_attn_metadata( + common_attn_metadata, draft_index=token_index + 1 + ) + ) # copy inputs to buffer for cudagraph self.input_ids[:batch_size] = input_ids @@ -655,6 +633,58 @@ class SpecDecodeBaseProposer: draft_token_ids = torch.stack(draft_token_ids_list, dim=1) return draft_token_ids + def _update_positions_dependent_metadata( + self, + positions: torch.Tensor, + common_attn_metadata, + batch_size: int, + input_batch_size: int, + block_size: int, + ) -> 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: + out_pos = self.mrope_positions[0, :batch_size] + elif self.uses_xdrope_dim > 0 and self.draft_uses_xdrope_dim > 0: + out_pos = self.xdrope_positions[0, :batch_size] + else: + out_pos = self.positions[:batch_size] + eagle_step_update_slot_mapping_and_metadata( + positions_1d=positions_1d, + block_table_tensor=common_attn_metadata.block_table_tensor, + seq_lens=common_attn_metadata.seq_lens, + block_size=block_size, + max_model_len=self.max_model_len, + out_clamped_positions=out_pos, + out_slot_mapping=self._slot_mapping_buffer[:input_batch_size], + input_batch_size=input_batch_size, + ) + common_attn_metadata.slot_mapping = self._slot_mapping_buffer[:batch_size] + if 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: + self.xdrope_positions[1:, :batch_size] = self.xdrope_positions[ + 0, :batch_size + ] + positions = self.xdrope_positions[0, :batch_size] + else: + positions = self.positions[:batch_size] + common_attn_metadata.max_seq_len = min( + common_attn_metadata.max_seq_len + 1, + self.max_model_len, + ) + + if common_attn_metadata._seq_lens_cpu is not None: + common_attn_metadata._seq_lens_cpu += 1 + if common_attn_metadata._num_computed_tokens_cpu is not None: + common_attn_metadata._num_computed_tokens_cpu += 1 + if common_attn_metadata.seq_lens_cpu_upper_bound is not None: + common_attn_metadata.seq_lens_cpu_upper_bound += 1 + + return positions + def set_inputs_first_pass( self, target_token_ids: torch.Tensor, diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 4975077c056..eef55957a08 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -169,6 +169,7 @@ from vllm.v1.spec_decode.dflash import DFlashProposer from vllm.v1.spec_decode.draft_model import DraftModelProposer from vllm.v1.spec_decode.eagle import EagleProposer from vllm.v1.spec_decode.extract_hidden_states import ExtractHiddenStatesProposer +from vllm.v1.spec_decode.gemma4 import Gemma4Proposer from vllm.v1.spec_decode.medusa import MedusaProposer from vllm.v1.spec_decode.metadata import SpecDecodeMetadata from vllm.v1.spec_decode.ngram_proposer_gpu import ( @@ -524,6 +525,7 @@ class GPUModelRunner( | DraftModelProposer | MedusaProposer | ExtractHiddenStatesProposer + | Gemma4Proposer ) if self.speculative_config.method == "ngram": from vllm.v1.spec_decode.ngram_proposer import NgramProposer @@ -552,6 +554,8 @@ class GPUModelRunner( self._ngram_pinned_val_buf = torch.zeros( self.max_num_reqs, dtype=torch.int32, pin_memory=True ) + elif self.speculative_config.use_gemma4_mtp(): + self.drafter = Gemma4Proposer(self.vllm_config, self.device, self) elif self.speculative_config.use_dflash(): self.drafter = DFlashProposer(self.vllm_config, self.device, self) self.use_aux_hidden_state_outputs = True @@ -2310,11 +2314,18 @@ class GPUModelRunner( cm.slot_mapping = slot_mappings[kv_cache_gid] if self.speculative_config and spec_decode_common_attn_metadata is None: - if isinstance(self.drafter, (EagleProposer, DFlashProposer)): + if isinstance( + self.drafter, (EagleProposer, DFlashProposer, Gemma4Proposer) + ): if self.drafter.kv_cache_gid == kv_cache_gid: spec_decode_common_attn_metadata = cm else: spec_decode_common_attn_metadata = cm + # Capture per-group block tables for multi-group proposers. + if self.speculative_config and isinstance(self.drafter, Gemma4Proposer): + self.drafter.set_per_group_block_table( + kv_cache_gid, cm.block_table_tensor + ) for attn_gid in range(len(self.attn_groups[kv_cache_gid])): if ubatch_slices is not None: @@ -4276,7 +4287,8 @@ class GPUModelRunner( EagleProposer | DFlashProposer | DraftModelProposer - | ExtractHiddenStatesProposer, + | ExtractHiddenStatesProposer + | Gemma4Proposer, ) sampled_token_ids = sampler_output.sampled_token_ids if input_fits_in_drafter: @@ -4672,7 +4684,8 @@ class GPUModelRunner( or spec_config.uses_draft_model() ): assert isinstance( - self.drafter, EagleProposer | DFlashProposer | DraftModelProposer + self.drafter, + EagleProposer | DFlashProposer | DraftModelProposer | Gemma4Proposer, ) if spec_config.disable_padded_drafter_batch: @@ -5594,7 +5607,8 @@ class GPUModelRunner( EagleProposer | DFlashProposer | DraftModelProposer - | ExtractHiddenStatesProposer, + | ExtractHiddenStatesProposer + | Gemma4Proposer, ) assert self.speculative_config is not None # Eagle currently only supports PIECEWISE cudagraphs. @@ -6395,7 +6409,8 @@ class GPUModelRunner( or self.speculative_config.uses_draft_model() ): assert isinstance( - self.drafter, EagleProposer | DFlashProposer | DraftModelProposer + self.drafter, + EagleProposer | DFlashProposer | DraftModelProposer | Gemma4Proposer, ) self.drafter.initialize_attn_backend(kv_cache_config, kernel_block_sizes) @@ -6448,7 +6463,10 @@ class GPUModelRunner( ): assert isinstance( self.drafter, - EagleProposer | DFlashProposer | ExtractHiddenStatesProposer, + EagleProposer + | DFlashProposer + | ExtractHiddenStatesProposer + | Gemma4Proposer, ) self.drafter.initialize_cudagraph_keys(cudagraph_mode)