FP8 MTP draft stack for vLLM — +6.6% decode on modelopt NVFP4 Qwen3.5-family checkpoints A vLLM contributor published an FP8 multi-token-prediction (MTP) draft stack for vLLM that yields a 6.6% decode throughput gain on modelopt NVFP4 Qwen3.5-family checkpoints. The change quantizes the previously BF16 draft stack (mtp.fc plus the MTP decoder layer's attention and MLP linears) to FP8 e4m3 with per-tensor scales and dynamic activation quantization, and adds an env-gated sliding window (MTP_DRAFT_WINDOW=8192) that bounds the draft's attention reads to a constant. Because the full-attention target still verifies every accepted token, the approach is lossless by construction and can only shift acceptance rate, not the output distribution. | | 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 |