# FP8 MTP draft stack for vLLM — +6.6% decode on modelopt NVFP4 Qwen3.5-family checkpoints

> Source: <https://gist.github.com/jaderinoo/0d1c0487073cef9d30ef4fd03499fa54>
> Published: 2026-09-15 17:52:06+00:00

|  | # 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 ( | 
|  | AutoWeightsLoader, | 
|  | 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", | 
|  | ) | 
|  | loader = AutoWeightsLoader(self) | 
|  | return loader.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 | 
|  |  | 
|  | loader = AutoWeightsLoader(self) | 
|  | loaded = loader.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() |
