| | # SPDX-License-Identifier: Apache-2.0 | | | # SPDX-FileCopyrightText: Copyright contributors to the vLLM project | | | """Inference-only Qwen3_5 MTP model.""" | | | | | | from collections.abc import Iterable | | | | | | import torch | | | from torch import nn | | | | | | from vllm._aiter_ops import rocm_aiter_ops | | | from vllm.compilation.decorators import support_torch_compile | | | from vllm.config import VllmConfig, get_current_vllm_config | | | from vllm.distributed import get_pp_group, tensor_model_parallel_all_gather | | | from vllm.logger import init_logger | | | from vllm.model_executor.layers.linear import ColumnParallelLinear | | | from vllm.model_executor.layers.logits_processor import LogitsProcessor | | | from vllm.model_executor.layers.vocab_parallel_embedding import ( | | | ParallelLMHead, | | | VocabParallelEmbedding, | | | ) | | | from vllm.model_executor.models.interfaces import LocalArgmaxMixin | | | from vllm.model_executor.models.qwen3_5 import ( | | | Qwen3_5DecoderLayer, | | | Qwen3_5Model, | | | Qwen3_5RMSNorm, | | | ) | | | from vllm.model_executor.models.qwen3_next import ( | | | QwenNextMixtureOfExperts, | | | _is_shared_expert_fse_compatible, | | | ) | | | from vllm.model_executor.models.utils import sequence_parallel_chunk | | | from vllm.sequence import IntermediateTensors | | | from vllm.transformers_utils.configs.qwen3_5 import Qwen3_5TextConfig | | | from vllm.transformers_utils.configs.qwen3_5_moe import Qwen3_5MoeTextConfig | | | | | | from .interfaces import ( | | | MultiModalEmbeddings, | | | SupportsMultiModal, | | | _require_is_multimodal, | | | ) | | | from .utils import ( | | | AutoWeights, | | | PPMissingLayer, | | | _merge_multimodal_embeddings, | | | make_empty_intermediate_tensors_factory, | | | maybe_fuse_shared_experts, | | | maybe_prefix, | | | ) | | | | | | logger = init_logger(name) | | | | | | # Windowed-MTP (arXiv 2607.21535): sliding window on the DRAFT's attention | | | # only. The MTP draft otherwise runs full attention over the whole KV cache | | | # at every draft step, so its read grows linearly with context; windowing | | | # bounds it to a constant. Lossless by construction — the full-attention | | | # target still verifies every accepted token, so only the proposals (and | | | # hence acceptance rate) can change, never the output distribution. | | | # Env-gated: MTP_DRAFT_WINDOW=8192 (absent/0 = off). | | | import os as _os | | | |
| | _draft_window = int(_os.environ.get("MTP_DRAFT_WINDOW", "0")) |
| | if _draft_window > 0: |
| | from vllm.model_executor.layers.attention import Attention as _Attention | | | | | | _orig_attention_init = _Attention.init | | | | | | def _mtp_windowed_attention_init(self, *args, **kwargs): | | | if "mtp.layers." in kwargs.get("prefix", ""): | | | kwargs.setdefault("per_layer_sliding_window", _draft_window) | | | _orig_attention_init(self, *args, *kwargs) | | | | | | _Attention.init = _mtp_windowed_attention_init | | | logger.info( | | | "Windowed-MTP: draft attention sliding window = %d tokens", | | | _draft_window, | | | ) | | | | | | # house: FP8 the BF16 draft stack (mtp.fc + the MTP decoder layer's | | | # attention/MLP linears). modelopt_fp4 checkpoints exclude mtp, so the whole | | | # draft stack is BF16 passthrough (~850 MiB/layer read per draft forward). | | | # With MTP_FP8=1 and an -mtpfp8 variant checkpoint (docker/quant/fp8_mtp.py) | | | # those linears load as FP8 e4m3 with per-tensor scales + dynamic activation | | | # quant. Draft proposals only — the target still verifies every token, so | | | # outputs are unchanged; only acceptance rate can shift. The 40960-row draft | | | # head stays BF16 (ParallelLMHead has no FP8 method). | | | _mtp_fp8 = _os.environ.get("MTP_FP8", "0") == "1" | | | | | | | | | def _mtp_fp8_config(): | | | from vllm.model_executor.layers.quantization.fp8 import Fp8Config | | | | | | return Fp8Config( | | | is_checkpoint_fp8_serialized=True, | | | activation_scheme="dynamic", | | | ) | | | | | | | | | class _Fp8DraftHeadMethod: | | | """quant_method for the vocab-truncated draft head under MTP_FP8. | | | ParallelLMHead has no FP8 path, so the BF16 head is converted in place |
| | after load (see Qwen3_5MTP.load_weights): weight stored column-major |
| | [K, N] e4m3 + per-tensor scale, dynamic per-token activation quant, |
| | CUTLASS scaled_mm — same recipe as Fp8LinearMethod. Draft proposals | | | only; the target still verifies every token.""" | | | | | | def apply(self, layer, x, bias=None): | | | from vllm import _custom_ops as ops | | | |
| | flat = x.reshape(-1, x.shape[-1]) |
| | qx, sx = ops.scaled_fp8_quant(flat, use_per_token_if_dynamic=True) |
| | return ops.cutlass_scaled_mm( | | | qx, layer.weight, sx, layer.weight_scale, x.dtype, bias | | | ) | | | | | | |
| | def _fp8_convert_draft_head(head: nn.Module) -> None: |
| | w = head.weight # [rows_per_rank, hidden] BF16 |
| | assert w.dtype == torch.bfloat16 and w.dim() == 2 |
| | fp8_max = torch.finfo(torch.float8_e4m3fn).max |
| | amax = w.float().abs().amax().clamp_min(1e-12) |
| | scale = (amax / fp8_max).reshape(1) |
| | qw = (w.float() / scale).clamp(-fp8_max, fp8_max).to(torch.float8_e4m3fn) |
| | head.weight = nn.Parameter(qw.t(), requires_grad=False) # col-major [K, N] |
| | head.weight_scale = scale | | | head.quant_method = _Fp8DraftHeadMethod() | | | logger.info( | | | "MTP_FP8: draft head (%d rows/rank) converted to FP8 e4m3", | | | qw.shape[0], | | | ) | | | | | | | | | @support_torch_compile( | | | dynamic_arg_dims={ | | | "input_ids": 0, | | | # positions is of shape (3, seq_len) if mrope is enabled for qwen2-vl, |
| | # otherwise (seq_len, ). |
| | "positions": -1, |
| | "intermediate_tensors": 0, | | | "inputs_embeds": 0, | | | "hidden_states": 0, | | | } | | | ) | | | class Qwen3_5MultiTokenPredictor(nn.Module): | | | hf_to_vllm_mapper = Qwen3_5Model.hf_to_vllm_mapper | | | |
| | def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): |
| | super().__init__() |
| | | | | model_config = vllm_config.model_config | | | quant_config = vllm_config.quant_config | | | | | | config: Qwen3_5TextConfig | Qwen3_5MoeTextConfig = model_config.hf_text_config | | | | | | self.config = config | | | | | | self.vocab_size = config.vocab_size | | | | | | self.mtp_start_layer_idx = config.num_hidden_layers | | | self.num_mtp_layers = getattr(config, "mtp_num_hidden_layers", 1) | | | | | | self.embed_tokens = VocabParallelEmbedding( | | | self.vocab_size, | | | config.hidden_size, | | | ) | | | | | | # syv patch: vocab-truncated draft head. If the checkpoint ships | | | # mtp_draft_vocab_ids.pt (built by build_draft_vocab.py) the drafter | | | # scores only those rows (mtp.draft_lm_head.*) instead of the full | | | # 248k-row lm_head; logits for all other ids are -inf. Speculative | | | # decoding stays exact, only the acceptance rate can change. | | | import os as _os | | | self.draft_lm_head = None | | | self.draft_vocab_ids = None | | | # house: model_config.model is an HF repo id here, not a local dir — | | | # resolve it to the cached snapshot so the ids file is found. | | | # snapshot_download refuses partially-fetched snapshots (vLLM never | | | # pulls README/LICENSE); try_to_load_from_cache just resolves the dir. | | | _model_dir = model_config.model | | | if not _os.path.isdir(_model_dir): | | | try: | | | from huggingface_hub import try_to_load_from_cache | | | | | | _cfg = try_to_load_from_cache(_model_dir, "config.json") |
| | if isinstance(_cfg, str): |
| | _model_dir = _os.path.dirname(_cfg) |
| | except Exception: | | | pass | | | _ids_path = _os.path.join(_model_dir, "mtp_draft_vocab_ids.pt") |
| | if _os.path.exists(_ids_path) and _os.environ.get("MTP_DRAFT_VOCAB", "1") != "0": |
| | _ids = torch.load(_ids_path, map_location="cpu") |
| | self.draft_vocab_ids = _ids | | | # house: our checkpoint stores the draft head as plain BF16 rows | | | # sliced from lm_head; modelopt_fp4 would otherwise try to load it | | | # as NVFP4 (same mtp.fc exclusion gap as below). Force unquantized. | | | _draft_quant = ( | | | None | | | if (quant_config and quant_config.get_name() == "modelopt_fp4") | | | else quant_config | | | ) | | | self.draft_lm_head = ParallelLMHead( | | | int(_ids.numel()), | | | config.hidden_size, | | | quant_config=_draft_quant, | | | prefix=maybe_prefix(prefix, "draft_lm_head"), | | | ) | | | logger.info("MTP drafter uses a %d-token draft head", int(_ids.numel())) | | | | | | # Workaround: mtp.fc is stored as BF16 in NVFP4 checkpoints but is | | | # missing from hf_quant_config.json exclude_modules. Force unquantized. | | | # Ref: https://github.com/vllm-project/vllm/pull/38650 | | | # Ref: https://github.com/NVIDIA/Model-Optimizer/pull/1124 |
| | if quant_config and quant_config.get_name() == "modelopt_fp4": |
| | fc_quant = _mtp_fp8_config() if _mtp_fp8 else None |
| | else: | | | fc_quant = quant_config | | | self.fc = ColumnParallelLinear( | | | self.config.hidden_size * 2, | | | self.config.hidden_size, | | | gather_output=True, | | | bias=False, | | | return_bias=False, | | | quant_config=fc_quant, | | | prefix=f"{prefix}.fc", | | | ) | | | | | | # GPTQ: quantized checkpoints may exclude MTP from quantization via | | | # quantization_config.dynamic with "-:pattern" entries. When detected, | | | # disable quantization for MTP layers so they use unquantized params. | | | original_quant = vllm_config.quant_config | | | if ( | | | _mtp_fp8 | | | and quant_config |
| | and quant_config.get_name() == "modelopt_fp4" |
| | ): |
| | # house: modelopt excludes mtp* (BF16 passthrough); run the MTP |
| | # decoder layer's linears as FP8 on the -mtpfp8 variant instead. |
| | vllm_config.quant_config = _mtp_fp8_config() |
| | elif quant_config and quant_config.get_name() not in ("modelopt_fp4",): |
| | hf_qc = getattr(model_config.hf_config, "quantization_config", None) |
| | if isinstance(hf_qc, dict): |
| | dynamic = hf_qc.get("dynamic", {}) |
| | if any(k.startswith("-:") and "mtp" in k for k in dynamic): |
| | vllm_config.quant_config = None | | | self.layers = torch.nn.ModuleList( | | | Qwen3_5DecoderLayer( | | | vllm_config, | | | layer_type="full_attention", | | | prefix=f"{prefix}.layers.{idx}", | | | ) | | | for idx in range(self.num_mtp_layers) | | | ) | | | vllm_config.quant_config = original_quant | | | self.make_empty_intermediate_tensors = make_empty_intermediate_tensors_factory( | | | ["hidden_states", "residual"], config.hidden_size | | | ) | | | self.norm = Qwen3_5RMSNorm(config.hidden_size, eps=config.rms_norm_eps) | | | self.pre_fc_norm_hidden = Qwen3_5RMSNorm( | | | config.hidden_size, eps=config.rms_norm_eps | | | ) | | | self.pre_fc_norm_embedding = Qwen3_5RMSNorm( | | | config.hidden_size, eps=config.rms_norm_eps | | | ) | | | | | | def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: | | | return self.embed_tokens(input_ids) | | | | | | 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, |
| | ) -> torch.Tensor: |
| | if get_pp_group().is_first_rank: |
| | if inputs_embeds is None: |
| | inputs_embeds = self.embed_input_ids(input_ids) |
| | assert hidden_states.shape[-1] == inputs_embeds.shape[-1] |
| | inputs_embeds = self.pre_fc_norm_embedding(inputs_embeds) | | | hidden_states = self.pre_fc_norm_hidden(hidden_states) |
| | hidden_states = torch.cat([inputs_embeds, hidden_states], dim=-1) |
| | hidden_states = self.fc(hidden_states) |
| | residual = None | | | else: | | | assert intermediate_tensors is not None | | | hidden_states = intermediate_tensors["hidden_states"] | | | residual = intermediate_tensors["residual"] | | | | | | current_step_idx = spec_step_idx % self.num_mtp_layers | | | mtp_layer = self.layers[current_step_idx] | | | if mtp_layer.use_attn_reduce_scatter_for_moe: | | | assert hidden_states.shape[0] == positions.shape[-1] | | | hidden_states = sequence_parallel_chunk(hidden_states) | | | assert residual is None | | | hidden_states, residual = mtp_layer( | | | positions=positions, | | | hidden_states=hidden_states, | | | residual=residual, | | | ) | | | | | | if not get_pp_group().is_last_rank: | | | return IntermediateTensors( | | | {"hidden_states": hidden_states, "residual": residual} | | | ) | | | | | | hidden_states, _ = self.norm(hidden_states, residual) | | | if mtp_layer.use_attn_reduce_scatter_for_moe: | | | hidden_states = tensor_model_parallel_all_gather(hidden_states, 0) | | | hidden_states = hidden_states[: positions.shape[-1]] | | | return hidden_states | | | | | | def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: | | | weights = maybe_fuse_shared_experts( | | | weights, | | | enabled=rocm_aiter_ops.is_fusion_moe_shared_experts_enabled() | | | and _is_shared_expert_fse_compatible( | | | get_current_vllm_config().quant_config | | | ), | | | n_routed_experts=getattr(self.config, "num_experts", 0), | | | n_shared_experts=1, | | | ckpt_prefix="mlp.shared_expert", | | | ) | | | = AutoWeights(self) | | | return .load_weights(weights, mapper=self.hf_to_vllm_mapper) | | | | | | | | | @support_torch_compile( | | | dynamic_arg_dims={ | | | "input_ids": 0, | | | # positions is of shape (3, seq_len) if mrope is enabled for qwen2-vl, |
| | # otherwise (seq_len, ). |
| | "positions": -1, |
| | "intermediate_tensors": 0, | | | "inputs_embeds": 0, | | | "hidden_states": 0, | | | } | | | ) | | | class Qwen3_5MTP(LocalArgmaxMixin, nn.Module, SupportsMultiModal): |
| | packed_modules_mapping = { |
| | "qkv_proj": [ |
| | "q_proj", | | | "k_proj", | | | "v_proj", | | | ], | | | "gate_up_proj": ["gate_proj", "up_proj"], | | | } | | | | | | def init(self, *, vllm_config: VllmConfig, prefix: str = ""): | | | config = vllm_config.model_config.hf_text_config | | | self.vllm_config = vllm_config | | | cache_config = vllm_config.cache_config | | | if cache_config.mamba_cache_mode == "all": | | | raise NotImplementedError( | | | "Qwen3_5MTP currently does not support 'all' prefix caching, " | | | "please use '--mamba-cache-mode=align' instead" | | | ) | | | | | | self.quant_config = vllm_config.quant_config | | | | | | super().init() | | | self.config = config | | | self.model = Qwen3_5MultiTokenPredictor( | | | vllm_config=vllm_config, prefix=maybe_prefix(prefix, "mtp") | | | ) | | | |
| | if get_pp_group().is_last_rank: |
| | self.lm_head = ParallelLMHead( |
| | config.vocab_size, | | | config.hidden_size, | | | quant_config=self.quant_config, | | | prefix=maybe_prefix(prefix, "lm_head"), | | | ) | | | if config.tie_word_embeddings: | | | self.lm_head = self.lm_head.tie_weights(self.model.embed_tokens) | | | else: | | | self.lm_head = PPMissingLayer() | | | | | | self.logits_processor = LogitsProcessor(config.vocab_size) | | | # syv patch: vocab-truncated draft head |
| | self.draft_logits_processor = ( |
| | LogitsProcessor(int(self.model.draft_vocab_ids.numel())) |
| | if getattr(self.model, "draft_lm_head", None) is not None | | | else None | | | ) | | | | | | def embed_input_ids( | | | self, | | | input_ids: torch.Tensor, | | | multimodal_embeddings: MultiModalEmbeddings | None = None, | | | *, | | | is_multimodal: torch.Tensor | None = None, | | | ) -> torch.Tensor: | | | inputs_embeds = self._embed_text_input_ids( | | | input_ids, | | | self.model.embed_input_ids, | | | is_multimodal=is_multimodal, | | | ) | | | | | | if multimodal_embeddings is None or len(multimodal_embeddings) == 0: | | | return inputs_embeds | | | | | | is_multimodal = _require_is_multimodal(is_multimodal) | | | | | | inputs_embeds = _merge_multimodal_embeddings( | | | inputs_embeds=inputs_embeds, | | | multimodal_embeddings=multimodal_embeddings, | | | is_multimodal=is_multimodal, | | | ) | | | | | | return inputs_embeds | | | | | | 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, | | | **kwargs: object, |
| | ): |
| | hidden_states = self.model( |
| | input_ids, positions, hidden_states, intermediate_tensors, inputs_embeds | | | ) | | | return hidden_states | | | | | | def compute_logits( | | | self, | | | hidden_states: torch.Tensor, |
| | spec_step_idx: int = 0, |
| | ) -> torch.Tensor \| None: |
| | # syv patch: vocab-truncated draft head | | | if self.draft_logits_processor is not None: | | | sub = self.draft_logits_processor(self.model.draft_lm_head, hidden_states) | | | if sub is None: | | | return None | | | ids = self.model.draft_vocab_ids |
| | if ids.device != sub.device: |
| | ids = ids.to(sub.device) |
| | self.model.draft_vocab_ids = ids |
| | full = sub.new_full((sub.shape[0], self.config.vocab_size), float("-inf")) |
| | full.index_copy_(1, ids, sub) |
| | return full | | | return self.logits_processor(self.lm_head, hidden_states) | | | |
| | def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: |
| | def remap_weight_names(weights): |
| | for name, weight in weights: | | | # syv patch: skip the truncated draft head when it is disabled | | | if "draft_lm_head" in name and self.model.draft_lm_head is None: | | | continue |
| | if name.startswith("mtp."): |
| | name = name.replace("mtp.", "model.") |
| | elif any(key in name for key in ["embed_tokens", "lm_head"]): |
| | if "embed_tokens" in name: | | | name = name.replace("language_model.", "") | | | else: | | | continue | | | yield name, weight | | | |
| | = AutoWeights(self) |
| | loaded = .load_weights(remap_weight_names(weights)) |
| | # house: MTP_FP8 also FP8s the BF16 draft head in place (no FP8 | | | # ParallelLMHead path in vLLM). Runs after load, before CUDA graphs. | | | if _mtp_fp8 and getattr(self.model, "draft_lm_head", None) is not None: | | | _fp8_convert_draft_head(self.model.draft_lm_head) | | | return loaded | | | | | | | | | class Qwen3_5MoeMTP(Qwen3_5MTP, QwenNextMixtureOfExperts): |
| | def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): |
| | super().__init__(vllm_config=vllm_config, prefix=prefix) |
| | self.set_moe_parameters() |