diff --git a/docs/models/supported_models.md b/docs/models/supported_models.md
index dec3975ed1a..0654a59caf1 100644
--- a/docs/models/supported_models.md
+++ b/docs/models/supported_models.md
@@ -633,6 +633,7 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen
| `SmolVLMForConditionalGeneration` | SmolVLM2 | T + I | `SmolVLM2-2.2B-Instruct` | ✅︎ | |
| `Step3VLForConditionalGeneration` | Step3-VL | T + I+ | `stepfun-ai/step3` | | ✅︎ |
| `StepVLForConditionalGeneration` | Step3-VL-10B | T + I+ | `stepfun-ai/Step3-VL-10B` | | ✅︎ |
+| `Step3p7ForConditionalGeneration` | Step-3.7-Flash | T + I+ | `stepfun-ai/Step-3.7-Flash` | | ✅︎ |
| `TarsierForConditionalGeneration` | Tarsier | T + IE+ | `omni-search/Tarsier-7b`, `omni-search/Tarsier-34b` | | ✅︎ |
| `Tarsier2ForConditionalGeneration`^ | Tarsier2 | T + IE+ + VE+ | `omni-research/Tarsier2-Recap-7b`, `omni-research/Tarsier2-7b-0115` | | ✅︎ |
| `UltravoxModel` | Ultravox | T + AE+ | `fixie-ai/ultravox-v0_5-llama-3_2-1b` | ✅︎ | ✅︎ |
diff --git a/tests/models/registry.py b/tests/models/registry.py
index 154eb3e479d..2b623152777 100644
--- a/tests/models/registry.py
+++ b/tests/models/registry.py
@@ -1368,6 +1368,9 @@ _MULTIMODAL_EXAMPLE_MODELS = {
"StepVLForConditionalGeneration": _HfExamplesInfo(
"stepfun-ai/Step3-VL-10B", trust_remote_code=True
),
+ "Step3p7ForConditionalGeneration": _HfExamplesInfo(
+ "stepfun-ai/Step-3.7-Flash", is_available_online=False, trust_remote_code=True
+ ),
"UltravoxModel": _HfExamplesInfo(
"fixie-ai/ultravox-v0_5-llama-3_2-1b",
trust_remote_code=True,
diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py
index 47d35f4ff4b..c76edaad0e7 100644
--- a/vllm/config/speculative.py
+++ b/vllm/config/speculative.py
@@ -485,7 +485,16 @@ class SpeculativeConfig:
{"n_predict": n_predict, "architectures": ["LongCatFlashMTPModel"]}
)
- if hf_config.model_type == "step3p5":
+ if hf_config.model_type in ("step3p5", "step3p7") or hf_config.architectures[
+ 0
+ ] in ("Step3p5ForCausalLM", "Step3p7ForConditionalGeneration"):
+ quantization_config = getattr(hf_config, "quantization_config", None)
+ hf_config = getattr(hf_config, "text_config", hf_config)
+ if (
+ quantization_config is not None
+ and getattr(hf_config, "quantization_config", None) is None
+ ):
+ hf_config.update({"quantization_config": quantization_config})
hf_config.model_type = "step3p5_mtp"
n_predict = getattr(hf_config, "num_nextn_predict_layers", 1)
hf_config.update({"n_predict": n_predict, "architectures": ["Step3p5MTP"]})
@@ -705,7 +714,11 @@ class SpeculativeConfig:
MTPModelTypes
):
self.method = "mtp"
- if self.num_speculative_tokens > 1:
+ if (
+ self.num_speculative_tokens > 1
+ and self.draft_model_config.hf_config.model_type
+ != "step3p5_mtp"
+ ):
logger.warning(
"Enabling num_speculative_tokens > 1 will run "
"multiple times of forward on same MTP layer"
@@ -1056,6 +1069,14 @@ class SpeculativeConfig:
== "gemma4_mtp"
)
+ def use_step3p5_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)
+ == "step3p5_mtp"
+ )
+
def use_eagle(self) -> bool:
return self.method in ("eagle", "eagle3", "mtp", "dflash")
diff --git a/vllm/model_executor/layers/fused_moe/experts/trtllm_fp8_moe.py b/vllm/model_executor/layers/fused_moe/experts/trtllm_fp8_moe.py
index 046d3e006af..b98b84cdc62 100644
--- a/vllm/model_executor/layers/fused_moe/experts/trtllm_fp8_moe.py
+++ b/vllm/model_executor/layers/fused_moe/experts/trtllm_fp8_moe.py
@@ -260,7 +260,7 @@ class TrtLlmFp8ExpertsMonolithic(TrtLlmFp8ExpertsBase, mk.FusedMoEExpertsMonolit
router_logits_dtype: torch.dtype | None,
routing_method: RoutingMethodType,
) -> bool:
- return True
+ return router_logits_dtype != torch.float32
@staticmethod
def _supports_routing_method(
diff --git a/vllm/model_executor/layers/fused_moe/experts/trtllm_nvfp4_moe.py b/vllm/model_executor/layers/fused_moe/experts/trtllm_nvfp4_moe.py
index ee56b50acd2..e4f292b7b1e 100644
--- a/vllm/model_executor/layers/fused_moe/experts/trtllm_nvfp4_moe.py
+++ b/vllm/model_executor/layers/fused_moe/experts/trtllm_nvfp4_moe.py
@@ -363,7 +363,7 @@ class TrtLlmNvFp4ExpertsMonolithic(
router_logits_dtype: torch.dtype | None,
routing_method: RoutingMethodType,
) -> bool:
- return True
+ return router_logits_dtype != torch.float32
def apply(
self,
diff --git a/vllm/model_executor/layers/fused_moe/layer.py b/vllm/model_executor/layers/fused_moe/layer.py
index b5d8f5fd016..d276db43a90 100644
--- a/vllm/model_executor/layers/fused_moe/layer.py
+++ b/vllm/model_executor/layers/fused_moe/layer.py
@@ -974,6 +974,19 @@ class FusedMoE(PluggableLayer):
# this is needed for compressed-tensors only
loaded_weight = loaded_weight.to(param.data.device)
+ # ModelOpt NVFP4 stores w13 input scales as two logical shards.
+ # The generic assignment below would broadcast w1/w3 into the
+ # whole expert row, so the second shard would overwrite the first.
+ if (
+ "ModelOpt" in quant_method_name
+ and param.data.ndim == 2
+ and shard_id in ("w1", "w3")
+ ):
+ scale_expert_id = global_expert_id if use_global_sf else expert_id
+ scale_shard_id = 0 if shard_id == "w1" else 1
+ param.data[scale_expert_id][scale_shard_id] = loaded_weight.reshape(())
+ return True if return_success else None
+
if (
"compressed" in quant_method_name.lower()
and param.data[expert_id] != 1
diff --git a/vllm/model_executor/layers/quantization/modelopt.py b/vllm/model_executor/layers/quantization/modelopt.py
index de224684788..eabaf62be78 100644
--- a/vllm/model_executor/layers/quantization/modelopt.py
+++ b/vllm/model_executor/layers/quantization/modelopt.py
@@ -192,7 +192,11 @@ class ModelOptQuantConfigBase(QuantizationConfig):
# exclude_modules config. But need to keep them for loading quantized
# checkpoints generated by older versions. Then check substring matching
# for patterns not caught by exact match
- if "vision_tower" in prefix or "vision_model" in prefix:
+ if (
+ "vision_tower" in prefix
+ or "vision_model" in prefix
+ or "vit_large_projector" in prefix
+ ):
return UnquantizedLinearMethod()
# now, the layer is quantized, handle it here
@@ -2340,6 +2344,14 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfigBase):
if key.startswith(prefix_dot):
return info["quant_algo"].upper()
+ # FusedMoE expert prefix is e.g. "...moe.experts", while ModelOpt's
+ # quantized_layers entries use "...moe.gate_proj" / "...moe.up_proj".
+ if prefix.endswith(".experts"):
+ parent_dot = prefix.rsplit(".experts", 1)[0] + "."
+ for key, info in self.quantized_layers.items():
+ if key.startswith(parent_dot):
+ return info["quant_algo"].upper()
+
return None
@staticmethod
diff --git a/vllm/model_executor/models/registry.py b/vllm/model_executor/models/registry.py
index 62c0149a3cf..94472d27e1c 100644
--- a/vllm/model_executor/models/registry.py
+++ b/vllm/model_executor/models/registry.py
@@ -567,6 +567,7 @@ _MULTIMODAL_MODELS = {
"SmolVLMForConditionalGeneration": ("smolvlm", "SmolVLMForConditionalGeneration"),
"StepVLForConditionalGeneration": ("step_vl", "StepVLForConditionalGeneration"),
"Step3VLForConditionalGeneration": ("step3_vl", "Step3VLForConditionalGeneration"),
+ "Step3p7ForConditionalGeneration": ("step3p7", "Step3p7ForConditionalGeneration"),
"TarsierForConditionalGeneration": ("tarsier", "TarsierForConditionalGeneration"),
"Tarsier2ForConditionalGeneration": (
"qwen2_vl",
diff --git a/vllm/model_executor/models/step3p5.py b/vllm/model_executor/models/step3p5.py
index cd73f6e26d9..c15cf18413b 100644
--- a/vllm/model_executor/models/step3p5.py
+++ b/vllm/model_executor/models/step3p5.py
@@ -637,6 +637,54 @@ class Step3p5Model(nn.Module):
(f".moe.experts.{base_layer}w13_weight", ".moe.gate_proj.weight", "w1"),
(f".moe.experts.{base_layer}w13_weight", ".moe.up_proj.weight", "w3"),
(f".moe.experts.{base_layer}w2_weight", ".moe.down_proj.weight", "w2"),
+ (
+ f".moe.experts.{base_layer}w13_weight_scale_2",
+ ".moe.gate_proj.weight_scale_2",
+ "w1",
+ ),
+ (
+ f".moe.experts.{base_layer}w13_weight_scale_2",
+ ".moe.up_proj.weight_scale_2",
+ "w3",
+ ),
+ (
+ f".moe.experts.{base_layer}w2_weight_scale_2",
+ ".moe.down_proj.weight_scale_2",
+ "w2",
+ ),
+ (
+ f".moe.experts.{base_layer}w13_weight_scale",
+ ".moe.gate_proj.weight_scale",
+ "w1",
+ ),
+ (
+ f".moe.experts.{base_layer}w13_weight_scale",
+ ".moe.up_proj.weight_scale",
+ "w3",
+ ),
+ (
+ f".moe.experts.{base_layer}w2_weight_scale",
+ ".moe.down_proj.weight_scale",
+ "w2",
+ ),
+ # Required due to the Step3 HF model's packed expert format:
+ # input scales are stored as moe.{gate,up,down}_proj.input_scale
+ # rather than the standard per-expert format handled generically.
+ (
+ f".moe.experts.{base_layer}w13_input_scale",
+ ".moe.gate_proj.input_scale",
+ "w1",
+ ),
+ (
+ f".moe.experts.{base_layer}w13_input_scale",
+ ".moe.up_proj.input_scale",
+ "w3",
+ ),
+ (
+ f".moe.experts.{base_layer}w2_input_scale",
+ ".moe.down_proj.input_scale",
+ "w2",
+ ),
]
# New per-expert format: .moe.experts.E.gate_proj.weight_packed [out, in]
@@ -756,7 +804,11 @@ class Step3p5Model(nn.Module):
# Per-tensor global scales (e.g. weight_global_scale)
# have shape [1] in compressed-tensors NVFP4 checkpoints.
# Expand to per-expert before the iteration loop.
- if (
+ if loaded_weight.ndim == 0:
+ loaded_weight = loaded_weight.unsqueeze(0).expand(
+ moe_expert_num
+ )
+ elif (
loaded_weight.shape[0] == 1
and loaded_weight.shape[0] != moe_expert_num
):
diff --git a/vllm/model_executor/models/step3p7.py b/vllm/model_executor/models/step3p7.py
new file mode 100644
index 00000000000..2a98ec53fc3
--- /dev/null
+++ b/vllm/model_executor/models/step3p7.py
@@ -0,0 +1,90 @@
+# SPDX-License-Identifier: Apache-2.0
+# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
+"""Inference-only Jurassic model."""
+
+import torch
+
+from vllm.config import VllmConfig
+from vllm.logger import init_logger
+from vllm.model_executor.layers.activation import get_act_fn
+from vllm.model_executor.layers.linear import ColumnParallelLinear
+
+from .step3_vl import Step3VLForConditionalGeneration
+from .step_vl import PerceptionEncoder
+from .utils import WeightsMapper, init_vllm_registered_model, maybe_prefix
+from .vision import run_dp_sharded_vision_model
+
+logger = init_logger(__name__)
+
+
+class Step3p7ForConditionalGeneration(Step3VLForConditionalGeneration):
+ hf_to_vllm_mapper = WeightsMapper(
+ orig_to_new_prefix={
+ "model.vision_model.": "vision_model.",
+ "model.vit_large_projector.": "vit_large_projector.",
+ "model.vit_large_projector": "vit_large_projector",
+ "model.language_model.": "language_model.model.",
+ "model.language_model": "language_model.model",
+ "model.": "language_model.model.",
+ "lm_head.": "language_model.lm_head.",
+ "lm_head": "language_model.lm_head",
+ },
+ orig_to_new_substr={
+ ".attn.in_proj_weight": ".attn.qkv_proj.weight",
+ ".attn.in_proj_bias": ".attn.qkv_proj.bias",
+ ".mlp.c_fc": ".mlp.fc1",
+ ".mlp.c_proj": ".mlp.fc2",
+ },
+ )
+
+ def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None:
+ super(Step3VLForConditionalGeneration, self).__init__()
+
+ config = vllm_config.model_config.hf_config
+ multimodal_config = vllm_config.model_config.multimodal_config
+ quant_config = vllm_config.quant_config
+
+ self.config = config
+ self.multimodal_config = multimodal_config
+ self.use_data_parallel = multimodal_config.mm_encoder_tp_mode == "data"
+
+ with self._mark_tower_model(vllm_config, "image"):
+ self.vision_model = PerceptionEncoder(
+ config.vision_config,
+ get_act_fn(config.vision_config.hidden_act),
+ quant_config=quant_config,
+ prefix=maybe_prefix(prefix, "vision_model"),
+ )
+ self.vit_large_projector = ColumnParallelLinear(
+ config.vision_config.width * 4,
+ config.text_config.hidden_size,
+ bias=config.projector_bias,
+ gather_output=True,
+ quant_config=quant_config,
+ prefix=maybe_prefix(prefix, "vit_large_projector"),
+ disable_tp=self.use_data_parallel,
+ )
+
+ with self._mark_language_model(vllm_config):
+ self.language_model = init_vllm_registered_model(
+ vllm_config=vllm_config,
+ hf_config=config.text_config,
+ prefix=maybe_prefix(prefix, "language_model"),
+ )
+
+ self.make_empty_intermediate_tensors = (
+ self.language_model.make_empty_intermediate_tensors
+ )
+
+ def _get_vision_model_output(
+ self, input_tensor: torch.Tensor | None
+ ) -> torch.Tensor | None:
+ if input_tensor is None:
+ return None
+ if self.use_data_parallel:
+ return run_dp_sharded_vision_model(input_tensor, self.vision_model)
+ return self.vision_model(input_tensor)
+
+ def _process_image_features(self, image_features: torch.Tensor) -> torch.Tensor:
+ image_features, _ = self.vit_large_projector(image_features)
+ return image_features
diff --git a/vllm/tokenizers/registry.py b/vllm/tokenizers/registry.py
index d72772ea00c..7578d3b43ab 100644
--- a/vllm/tokenizers/registry.py
+++ b/vllm/tokenizers/registry.py
@@ -38,7 +38,7 @@ logger = init_logger(__name__)
# temporary workaround and better long term solutions are:
# - Add model type to MODELS_WITH_INCORRECT_HUB_TOKENIZER_CLASS in transformers (better)
# - Fix tokenizer_class on the hub for the affected models (best)
-_MODEL_TYPES_WITH_INCORRECT_TOKENIZER_CLASS: set[str] = {"step3_vl"}
+_MODEL_TYPES_WITH_INCORRECT_TOKENIZER_CLASS: set[str] = {"step3_vl", "step3p7"}
_VLLM_TOKENIZERS = {
"deepseek_v32": ("deepseek_v32", "DeepseekV32Tokenizer"),
@@ -249,6 +249,10 @@ def get_tokenizer(
tokenizer_cls_ = tokenizer_cls
tokenizer = tokenizer_cls_.from_pretrained(tokenizer_name, *args, **kwargs)
+ if model_type in _MODEL_TYPES_WITH_INCORRECT_TOKENIZER_CLASS:
+ from vllm.tokenizers.hf import get_cached_tokenizer
+
+ tokenizer = get_cached_tokenizer(tokenizer)
if not tokenizer.is_fast:
logger.warning(
"Using a slow tokenizer. This might cause a significant "
diff --git a/vllm/transformers_utils/model_arch_config_convertor.py b/vllm/transformers_utils/model_arch_config_convertor.py
index 35fa1313d1e..85452197535 100644
--- a/vllm/transformers_utils/model_arch_config_convertor.py
+++ b/vllm/transformers_utils/model_arch_config_convertor.py
@@ -502,6 +502,11 @@ class Qwen3_5MTPModelArchConfigConvertor(ModelArchConfigConvertorBase):
return getattr(self.hf_text_config, "mtp_num_hidden_layers", 0)
+class Step3p5MTPModelArchConfigConvertor(ModelArchConfigConvertorBase):
+ def get_num_hidden_layers(self) -> int:
+ return getattr(self.hf_text_config, "num_nextn_predict_layers", 0)
+
+
class PanguUltraMoeMTPModelArchConfigConvertor(ModelArchConfigConvertorBase):
def get_num_hidden_layers(self) -> int:
return getattr(self.hf_text_config, "num_nextn_predict_layers", 0)
@@ -543,31 +548,32 @@ class Gemma4ModelArchConfigConvertor(ModelArchConfigConvertorBase):
# hf_config.model_type -> convertor class
MODEL_ARCH_CONFIG_CONVERTORS = {
"cohere_asr": CohereAsrModelArchConfigConvertor,
- "mamba": MambaModelArchConfigConvertor,
- "falcon_mamba": MambaModelArchConfigConvertor,
- "timm_wrapper": TerratorchModelArchConfigConvertor,
- "medusa": MedusaModelArchConfigConvertor,
- "zamba2": Zamba2ModelArchConfigConvertor,
- "mpt": MPTModelArchConfigConvertor,
"dbrx": DbrxModelArchConfigConvertor,
- "falcon": FalconModelArchConfigConvertor,
- "gemma4": Gemma4ModelArchConfigConvertor,
- "gemma4_text": Gemma4ModelArchConfigConvertor,
- "gemma4_mtp": Gemma4MTPModelArchConfigConvertor,
- "RefinedWeb": FalconModelArchConfigConvertor,
- "RefinedWebModel": FalconModelArchConfigConvertor,
- "nemotron-nas": NemotronNasModelArchConfigConvertor,
"deepseek_mtp": DeepSeekMTPModelArchConfigConvertor,
- "qwen3_next_mtp": Qwen3NextMTPModelArchConfigConvertor,
- "qwen3_5_mtp": Qwen3_5MTPModelArchConfigConvertor,
+ "ernie_mtp": ErnieMTPModelArchConfigConvertor,
+ "falcon": FalconModelArchConfigConvertor,
+ "falcon_mamba": MambaModelArchConfigConvertor,
+ "gemma4": Gemma4ModelArchConfigConvertor,
+ "gemma4_mtp": Gemma4MTPModelArchConfigConvertor,
+ "gemma4_text": Gemma4ModelArchConfigConvertor,
+ "glm4_moe_mtp": GLM4MoeMTPModelArchConfigConvertor,
+ "glm_ocr_mtp": GLM4MoeMTPModelArchConfigConvertor,
+ "longcat_flash_mtp": LongCatFlashMTPModelArchConfigConvertor,
+ "mamba": MambaModelArchConfigConvertor,
+ "medusa": MedusaModelArchConfigConvertor,
"mimo_mtp": MimoMTPModelArchConfigConvertor,
"mimo_v2": MimoV2ModelArchConfigConvertor,
"mimo_v2_flash": MimoV2ModelArchConfigConvertor,
"mimo_v2_mtp": MimoV2MTPModelArchConfigConvertor,
"mimo_v2_omni_mtp": MimoV2MTPModelArchConfigConvertor,
- "glm4_moe_mtp": GLM4MoeMTPModelArchConfigConvertor,
- "glm_ocr_mtp": GLM4MoeMTPModelArchConfigConvertor,
- "ernie_mtp": ErnieMTPModelArchConfigConvertor,
+ "mpt": MPTModelArchConfigConvertor,
+ "nemotron-nas": NemotronNasModelArchConfigConvertor,
"pangu_ultra_moe_mtp": PanguUltraMoeMTPModelArchConfigConvertor,
- "longcat_flash_mtp": LongCatFlashMTPModelArchConfigConvertor,
+ "qwen3_5_mtp": Qwen3_5MTPModelArchConfigConvertor,
+ "qwen3_next_mtp": Qwen3NextMTPModelArchConfigConvertor,
+ "RefinedWeb": FalconModelArchConfigConvertor,
+ "RefinedWebModel": FalconModelArchConfigConvertor,
+ "step3p5_mtp": Step3p5MTPModelArchConfigConvertor,
+ "timm_wrapper": TerratorchModelArchConfigConvertor,
+ "zamba2": Zamba2ModelArchConfigConvertor,
}
diff --git a/vllm/v1/spec_decode/llm_base_proposer.py b/vllm/v1/spec_decode/llm_base_proposer.py
index 9979a051727..9a0b537175b 100644
--- a/vllm/v1/spec_decode/llm_base_proposer.py
+++ b/vllm/v1/spec_decode/llm_base_proposer.py
@@ -1232,6 +1232,7 @@ class SpecDecodeBaseProposer:
"Qwen3VLForConditionalGeneration",
"Qwen3VLMoeForConditionalGeneration",
"Gemma4ForConditionalGeneration",
+ "Step3p7ForConditionalGeneration",
]:
self.model.config.image_token_index = target_model.config.image_token_id
elif self.get_model_name(target_model) == "PixtralForConditionalGeneration":
diff --git a/vllm/v1/spec_decode/step3p5.py b/vllm/v1/spec_decode/step3p5.py
new file mode 100644
index 00000000000..ccca17a3188
--- /dev/null
+++ b/vllm/v1/spec_decode/step3p5.py
@@ -0,0 +1,459 @@
+# SPDX-License-Identifier: Apache-2.0
+# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
+
+from copy import copy
+
+import torch
+
+from vllm.config import VllmConfig, get_layers_from_vllm_config, replace
+from vllm.forward_context import set_forward_context
+from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
+from vllm.model_executor.models.utils import get_draft_quant_config
+from vllm.v1.attention.backend import CommonAttentionMetadata
+from vllm.v1.kv_cache_interface import (
+ KVCacheConfig,
+ KVCacheSpec,
+ UniformTypeKVCacheSpecs,
+)
+from vllm.v1.sample.metadata import SamplingMetadata
+from vllm.v1.spec_decode.eagle import EagleProposer
+from vllm.v1.spec_decode.utils import PADDING_SLOT_ID
+from vllm.v1.worker.utils import AttentionGroup
+
+
+class Step3p5MTPProposer(EagleProposer):
+ """Step3.5 MTP proposer with per-layer draft-step selection."""
+
+ def __init__(
+ self,
+ vllm_config: VllmConfig,
+ device: torch.device,
+ runner=None,
+ ):
+ super().__init__(vllm_config, device, runner)
+ self._per_group_block_tables: dict[int, torch.Tensor] = {}
+ self._per_group_slot_mappings: dict[int, torch.Tensor] = {}
+ # Slot-mapping buffers for non-primary KV cache groups (the primary
+ # group reuses self._slot_mapping_buffer from the base class).
+ self._per_group_slot_mapping_buffers: dict[int, torch.Tensor] = {}
+
+ def set_per_group_attn_metadata(
+ self,
+ gid: int,
+ block_table: torch.Tensor,
+ slot_mapping: torch.Tensor,
+ ) -> None:
+ self._per_group_block_tables[gid] = block_table
+ self._per_group_slot_mappings[gid] = slot_mapping
+
+ def _slot_mapping_buffer_for(self, gid: int) -> torch.Tensor:
+ if gid == self.kv_cache_gid:
+ return self._slot_mapping_buffer
+ buf = self._per_group_slot_mapping_buffers.get(gid)
+ if buf is None:
+ buf = torch.zeros(self.max_positions, dtype=torch.int64, device=self.device)
+ self._per_group_slot_mapping_buffers[gid] = buf
+ return buf
+
+ def _get_slot_mapping(
+ self,
+ num_tokens: int,
+ slot_mapping: torch.Tensor | None = None,
+ ) -> dict[str, torch.Tensor]:
+ """Per-layer slot_mapping with one buffer per KV cache group."""
+ per_layer: dict[str, torch.Tensor] = {}
+ for attn_group in self.draft_attn_groups:
+ gid = attn_group.kv_cache_group_id
+ buf = self._slot_mapping_buffer_for(gid)
+ source = self._per_group_slot_mappings.get(gid, slot_mapping)
+ if source is not None and buf.data_ptr() != source.data_ptr():
+ n = source.shape[0]
+ buf[:n].copy_(source)
+ if num_tokens > n:
+ buf[n:num_tokens].fill_(PADDING_SLOT_ID)
+ view = buf[:num_tokens]
+ for layer_name in attn_group.layer_names:
+ per_layer[layer_name] = view
+ return per_layer
+
+ def _update_positions_dependent_metadata(
+ self,
+ positions: torch.Tensor,
+ common_attn_metadata: CommonAttentionMetadata,
+ batch_size: int,
+ input_batch_size: int,
+ block_size: int,
+ ) -> torch.Tensor:
+ old_positions_1d = positions[0] if self.uses_mrope else positions
+ positions = super()._update_positions_dependent_metadata(
+ positions,
+ common_attn_metadata,
+ batch_size,
+ input_batch_size,
+ block_size,
+ )
+ # Parent already produced slot_mapping for the primary gid.
+ self._per_group_slot_mappings[self.kv_cache_gid] = (
+ common_attn_metadata.slot_mapping
+ )
+ # Recompute slot_mapping for the remaining gids using their own block tables.
+ new_positions_1d = positions[0] if self.uses_mrope else positions
+ exceeds = old_positions_1d + 1 >= self.max_model_len
+ for attn_group in self.draft_attn_groups:
+ gid = attn_group.kv_cache_group_id
+ if gid == self.kv_cache_gid:
+ continue
+ block_table = self._per_group_block_tables.get(gid)
+ if block_table is None:
+ continue
+ n_blocks = block_table.shape[1]
+ bn = (new_positions_1d // block_size).clamp(max=n_blocks - 1).to(torch.long)
+ block_ids = block_table[:batch_size].gather(1, bn.unsqueeze(1)).squeeze(1)
+ sm = block_ids * block_size + (new_positions_1d % block_size)
+ sm.masked_fill_(exceeds, PADDING_SLOT_ID)
+ buf = self._slot_mapping_buffer_for(gid)
+ buf[:batch_size].copy_(sm)
+ if input_batch_size > batch_size:
+ buf[batch_size:input_batch_size].fill_(PADDING_SLOT_ID)
+ self._per_group_slot_mappings[gid] = buf[:batch_size]
+ return positions
+
+ def build_per_group_and_layer_attn_metadata(
+ self,
+ common_attn_metadata: CommonAttentionMetadata,
+ draft_index: int = 0,
+ ) -> tuple[list[object], dict[str, object]]:
+ per_group_attn_metadata: list[object] = []
+ per_layer_attn_metadata: dict[str, object] = {}
+ # The proposer always works in unpadded shape. Per-group block tables
+ # registered via set_per_group_attn_metadata are stored at the model
+ # runner's padded shape; slice them to match cm's num_reqs.
+ num_reqs = common_attn_metadata.num_reqs
+ num_actual_tokens = common_attn_metadata.num_actual_tokens
+ 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][:num_reqs]
+ if gid in self._per_group_slot_mappings:
+ sm = self._per_group_slot_mappings[gid]
+ if sm.shape[0] >= num_actual_tokens:
+ sm = sm[:num_actual_tokens]
+ cm.slot_mapping = sm
+ 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 _maybe_share_lm_head(self, target_language_model: torch.nn.Module) -> None:
+ """Step3.5 MTP uses the lm_head stored in each MTP layer."""
+
+ # The base MTP path shares target lm_head into shared_head.head.
+ # Step3.5 checkpoints carry per-MTP-layer shared_head weights.
+ return
+
+ def _create_draft_vllm_config(self) -> VllmConfig:
+ base = super()._create_draft_vllm_config()
+ return replace(
+ base,
+ model_config=self.draft_model_config,
+ quant_config=get_draft_quant_config(base),
+ )
+
+ def validate_same_kv_cache_group(self, kv_cache_config: KVCacheConfig) -> None:
+ """Step3.5 MTP draft layers may span multiple KV cache groups."""
+ return
+
+ def initialize_attn_backend(
+ self,
+ kv_cache_config: KVCacheConfig,
+ kernel_block_sizes: list[int] | None = None,
+ ) -> None:
+ 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 layer_name in group.layer_names:
+ layer_to_gid[layer_name] = gid
+ if isinstance(group_spec, UniformTypeKVCacheSpecs):
+ if layer_name in group_spec.kv_cache_specs:
+ layer_to_spec[layer_name] = group_spec.kv_cache_specs[
+ layer_name
+ ]
+ else:
+ target_layer_name = getattr(
+ all_attn_layers.get(layer_name),
+ "kv_sharing_target_layer_name",
+ None,
+ )
+ if (
+ target_layer_name
+ and target_layer_name in group_spec.kv_cache_specs
+ ):
+ layer_to_spec[layer_name] = group_spec.kv_cache_specs[
+ target_layer_name
+ ]
+ else:
+ layer_to_spec[layer_name] = group_spec
+ else:
+ layer_to_spec[layer_name] = group_spec
+
+ attention_groups: dict[tuple[tuple[str, str], int], AttentionGroup] = {}
+ for layer_name in sorted(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(), gid)
+
+ 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
+
+ def _sample_draft_tokens_for_step(
+ self,
+ hidden_states: torch.Tensor,
+ sampling_metadata: SamplingMetadata,
+ spec_step_idx: int,
+ ) -> tuple[torch.Tensor, torch.Tensor | None]:
+ if not self._enable_probabilistic_draft_probs or sampling_metadata.all_greedy:
+ if self.use_local_argmax_reduction:
+ return self.model.get_top_tokens(hidden_states), None
+ logits = self.model.compute_logits(
+ hidden_states, spec_step_idx=spec_step_idx
+ )
+ return logits.argmax(dim=-1), None
+
+ logits = self.model.compute_logits(hidden_states, spec_step_idx=spec_step_idx)
+ return self._sample_from_logits(logits, sampling_metadata)
+
+ def propose(
+ self,
+ target_token_ids: torch.Tensor,
+ target_positions: torch.Tensor,
+ target_hidden_states: torch.Tensor,
+ next_token_ids: torch.Tensor,
+ token_indices_to_sample: torch.Tensor | None,
+ common_attn_metadata: CommonAttentionMetadata,
+ sampling_metadata: SamplingMetadata,
+ mm_embed_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None,
+ num_rejected_tokens_gpu: torch.Tensor | None = None,
+ slot_mappings: dict[str, torch.Tensor]
+ | list[dict[str, torch.Tensor]]
+ | None = None,
+ ) -> torch.Tensor:
+ self._last_draft_probs = None
+ batch_size = common_attn_metadata.batch_size()
+
+ num_tokens, token_indices_to_sample, common_attn_metadata = (
+ self.set_inputs_first_pass(
+ target_token_ids=target_token_ids,
+ next_token_ids=next_token_ids,
+ target_positions=target_positions,
+ target_hidden_states=target_hidden_states,
+ token_indices_to_sample=token_indices_to_sample,
+ cad=common_attn_metadata,
+ num_rejected_tokens_gpu=num_rejected_tokens_gpu,
+ )
+ )
+
+ per_group_attn_metadata, per_layer_attn_metadata = (
+ self.build_per_group_and_layer_attn_metadata(common_attn_metadata)
+ )
+
+ cudagraph_runtime_mode, num_input_tokens, num_tokens_across_dp = (
+ self._determine_batch_execution_and_padding(num_tokens)
+ )
+
+ model_kwargs, slot_mapping_size = self.build_model_inputs_first_pass(
+ num_tokens, num_input_tokens, mm_embed_inputs
+ )
+ model_kwargs["spec_step_idx"] = 0
+
+ with set_forward_context(
+ per_layer_attn_metadata,
+ self.vllm_config,
+ num_tokens=num_input_tokens,
+ num_tokens_across_dp=num_tokens_across_dp,
+ cudagraph_runtime_mode=cudagraph_runtime_mode,
+ slot_mapping=self._get_slot_mapping(
+ slot_mapping_size, common_attn_metadata.slot_mapping
+ ),
+ ):
+ ret_hidden_states = self.model(**model_kwargs)
+ if not self.model_returns_tuple():
+ last_hidden_states = ret_hidden_states
+ hidden_states = last_hidden_states
+ else:
+ last_hidden_states, hidden_states = ret_hidden_states
+
+ sample_hidden_states = last_hidden_states[token_indices_to_sample]
+
+ if self.num_speculative_tokens == 1 or self.parallel_drafting:
+ draft_token_ids, draft_probs = self._sample_draft_tokens_for_step(
+ sample_hidden_states, sampling_metadata, spec_step_idx=0
+ )
+ if draft_probs is not None:
+ self._last_draft_probs = draft_probs.view(
+ -1, self.num_speculative_tokens, draft_probs.shape[-1]
+ ).contiguous()
+ return draft_token_ids.view(-1, self.num_speculative_tokens)
+
+ if self.uses_mrope:
+ positions = self.mrope_positions[:, token_indices_to_sample]
+ else:
+ positions = self.positions[token_indices_to_sample]
+ hidden_states = hidden_states[token_indices_to_sample]
+
+ if self.constant_draft_positions:
+ self.positions[:batch_size] = positions
+
+ draft_token_ids, draft_probs = self._sample_draft_tokens_for_step(
+ sample_hidden_states, sampling_metadata, spec_step_idx=0
+ )
+ draft_probs_list = None if draft_probs is None else [draft_probs]
+
+ if self.allowed_attn_types is not None:
+ for group_md in per_group_attn_metadata:
+ if not isinstance(group_md, self.allowed_attn_types):
+ raise ValueError(
+ f"Unsupported attention metadata type for speculative "
+ "decoding with num_speculative_tokens > 1: "
+ f"{type(group_md)}. Supported types are: "
+ f"{self.allowed_attn_types}"
+ )
+
+ draft_token_ids_list = [draft_token_ids]
+
+ cudagraph_runtime_mode, input_batch_size, batch_size_across_dp = (
+ self._determine_batch_execution_and_padding(batch_size)
+ )
+
+ common_attn_metadata.num_actual_tokens = batch_size
+ common_attn_metadata.max_query_len = 1
+ common_attn_metadata.query_start_loc = self.arange[: batch_size + 1]
+ common_attn_metadata.query_start_loc_cpu = torch.from_numpy(
+ self.token_arange_np[: batch_size + 1]
+ ).clone()
+
+ if self.num_speculative_tokens > 1 and num_rejected_tokens_gpu is not None:
+ common_attn_metadata.seq_lens -= num_rejected_tokens_gpu
+ common_attn_metadata._seq_lens_cpu = None
+ common_attn_metadata._num_computed_tokens_cpu = None
+
+ block_size = self.block_size
+ assert block_size > 0, "block_size has not been initialized."
+ for token_index in range(self.num_speculative_tokens - 1):
+ spec_step_idx = token_index + 1
+ input_ids = draft_token_ids_list[-1].int()
+
+ if not self.constant_draft_positions:
+ positions = self._update_positions_dependent_metadata(
+ positions,
+ common_attn_metadata,
+ batch_size,
+ input_batch_size,
+ block_size,
+ )
+
+ 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=spec_step_idx
+ )
+ )
+
+ self.input_ids[:batch_size] = input_ids
+ self.hidden_states[:batch_size] = hidden_states
+ if self.supports_mm_inputs:
+ self.inputs_embeds[:batch_size] = self.model.embed_input_ids(input_ids)
+
+ input_ids = None
+ inputs_embeds = self.inputs_embeds[:input_batch_size]
+ else:
+ input_ids = self.input_ids[:input_batch_size]
+ inputs_embeds = None
+
+ model_kwargs = {
+ "input_ids": input_ids,
+ "positions": self._get_positions(input_batch_size),
+ "inputs_embeds": inputs_embeds,
+ "spec_step_idx": spec_step_idx,
+ }
+ if self.pass_hidden_states_to_model:
+ model_kwargs["hidden_states"] = self.hidden_states[:input_batch_size]
+
+ with set_forward_context(
+ per_layer_attn_metadata,
+ self.vllm_config,
+ num_tokens=input_batch_size,
+ num_tokens_across_dp=batch_size_across_dp,
+ cudagraph_runtime_mode=cudagraph_runtime_mode,
+ slot_mapping=self._get_slot_mapping(input_batch_size),
+ ):
+ ret_hidden_states = self.model(**model_kwargs)
+ if not self.model_returns_tuple():
+ last_hidden_states = ret_hidden_states
+ hidden_states = ret_hidden_states
+ else:
+ last_hidden_states, hidden_states = ret_hidden_states
+
+ hidden_states = hidden_states[:batch_size]
+ draft_token_ids, draft_probs = self._sample_draft_tokens_for_step(
+ last_hidden_states[:batch_size],
+ sampling_metadata,
+ spec_step_idx=spec_step_idx,
+ )
+ if draft_probs is not None:
+ assert draft_probs_list is not None
+ draft_probs_list.append(draft_probs)
+ draft_token_ids_list.append(draft_token_ids)
+
+ draft_token_ids = torch.stack(draft_token_ids_list, dim=1)
+ if draft_probs_list is not None:
+ self._last_draft_probs = torch.stack(draft_probs_list, dim=1).contiguous()
+ return draft_token_ids
diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py
index 779d73921e8..f82d2224a41 100644
--- a/vllm/v1/worker/gpu_model_runner.py
+++ b/vllm/v1/worker/gpu_model_runner.py
@@ -186,6 +186,7 @@ from vllm.v1.spec_decode.ngram_proposer_gpu import (
update_ngram_gpu_tensors_incremental,
update_scheduler_for_invalid_drafts,
)
+from vllm.v1.spec_decode.step3p5 import Step3p5MTPProposer
from vllm.v1.spec_decode.suffix_decoding import SuffixDecodingProposer
from vllm.v1.spec_decode.utils import update_num_computed_tokens_for_batch_change
from vllm.v1.structured_output.utils import apply_grammar_bitmask
@@ -547,6 +548,7 @@ class GPUModelRunner(
| MedusaProposer
| ExtractHiddenStatesProposer
| Gemma4Proposer
+ | Step3p5MTPProposer
)
if self.speculative_config.method == "custom_class":
self.drafter = create_custom_proposer( # type: ignore[assignment]
@@ -581,6 +583,8 @@ class GPUModelRunner(
)
elif self.speculative_config.use_gemma4_mtp():
self.drafter = Gemma4Proposer(self.vllm_config, self.device, self)
+ elif self.speculative_config.use_step3p5_mtp():
+ self.drafter = Step3p5MTPProposer(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
@@ -2428,7 +2432,11 @@ class GPUModelRunner(
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):
+ if self.speculative_config and isinstance(self.drafter, Step3p5MTPProposer):
+ self.drafter.set_per_group_attn_metadata(
+ kv_cache_gid, cm.block_table_tensor, cm.slot_mapping
+ )
+ elif self.speculative_config and isinstance(self.drafter, Gemma4Proposer):
self.drafter.set_per_group_block_table(
kv_cache_gid, cm.block_table_tensor
)