From bae8fa23a297d62e6a76754cfb2093c56bb87f03 Mon Sep 17 00:00:00 2001 From: Kai Xu Date: Wed, 23 Sep 2026 13:24:00 -0700 Subject: [PATCH] Add vLLM GDN/KDA state-only fake quantization Signed-off-by: Kai Xu --- CHANGELOG.rst | 2 + examples/llm_qat/linear_attention/README.md | 6 + examples/vllm_serve/README.md | 77 +++++ examples/vllm_serve/fakequant_worker.py | 26 +- .../linear_attention_state_int8.yaml | 12 + examples/vllm_serve/vllm_reload_utils.py | 45 ++- .../torch/quantization/plugins/__init__.py | 1 + .../plugins/vllm_linear_attention.py | 165 +++++++++ .../test_vllm_linear_attention.py | 315 ++++++++++++++++++ 9 files changed, 638 insertions(+), 11 deletions(-) create mode 100644 examples/vllm_serve/linear_attention_state_int8.yaml create mode 100644 modelopt/torch/quantization/plugins/vllm_linear_attention.py create mode 100644 tests/gpu_vllm/torch/quantization/test_vllm_linear_attention.py diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 60d905778b0..0a751f6c2d5 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -13,6 +13,8 @@ Changelog *Quantization* +- Add experimental FP8/INT8 recurrent-state fake quantization before native vLLM GDN/KDA prefill and decode. Use the state-only recipe with the documented vLLM 0.15.x eager runtime. + - Add experimental GDN/KDA decode-aware QAT with FP8 or INT8 recurrent states, optional INT8 value-axis Hadamard encoding, KDA decay rounding, and encoded-update replay. Supply explicit prefix lengths through the training phase context; prefix solves remain exact. - Add ``layerwise.export_dir``: layerwise calibration writes each decoder layer to its own quantized checkpoint shard as it finishes, so no separate ``export_hf_checkpoint()`` pass is needed and, with ``layerwise.checkpoint_dir``, an interrupted run resumes without redoing finished layers. Calibration writes the layer shards; ``finalize()`` on the exporter left on the model adds the tail shard, the index and the config artifacts, and the checkpoint does not load until it runs. ``examples/hf_ptq`` does this for you. Supports FP8 and NVFP4 on single-process models, resident or offloaded, including multimodal models and models with MTP layers; other formats and placements raise ``NotImplementedError`` before calibration starts. diff --git a/examples/llm_qat/linear_attention/README.md b/examples/llm_qat/linear_attention/README.md index ad0e116c1c3..43b6c970c39 100644 --- a/examples/llm_qat/linear_attention/README.md +++ b/examples/llm_qat/linear_attention/README.md @@ -305,3 +305,9 @@ has no prefill GEMM QDQ or approximate inverse. FLA KDA requires `use_cache=Fals serving cache objects are rejected. ModelOpt saves execution policies, while per-batch prefix lengths must be supplied again during training. Distributed decode training and model-quality recovery require separate qualification. + +## State-only serving with vLLM + +See the [vLLM example](../../vllm_serve/README.md#linear-attention-state-quantization) +for incoming-state QDQ before native prefill/decode. Its invocation boundaries +differ from the token-write and ReplaySSM policies used during training. diff --git a/examples/vllm_serve/README.md b/examples/vllm_serve/README.md index 9c4eec609b7..e12e429f17d 100644 --- a/examples/vllm_serve/README.md +++ b/examples/vllm_serve/README.md @@ -7,6 +7,10 @@ Compared with realquant, fakequant is 2-5x slower, but doesn't require dedicated The general fakequant example is tested with vLLM 0.9.0, 0.19.1, 0.26.0, and 0.28.0. The compact NVFP4 attention worker documented below requires vLLM 0.15.0 or newer. +For GDN/KDA recurrent-state QDQ before native prefill and decode, see the +[linear-attention example](#linear-attention-state-quantization), which +documents the initial vLLM 0.15.x runtime requirements and state-only INT8 recipe. + ## Prepare environment Use the Dockerfile to build an environment with vLLM 0.28.0: @@ -279,6 +283,79 @@ Supported configurations are regular decoder self-attention with FlashInfer or F Unsupported features are sliding window, ALiBi, softcap, sinks, FP8 KV cache, cross/encoder/MLA attention, KV sharing or transfer, prefix caching, speculative decoding, DBO/ubatching, and `FULL` mixed/prefill CUDA graphs. +## Linear-attention state quantization + +The existing `FakeQuantWorker` adds a ModelOpt `TensorQuantizer` immediately before +native vLLM prefill and decode for `Qwen3NextGatedDeltaNet` and +`KimiDeltaAttention`. Native attention kernels, projections, convolution, gates, +outputs, and cache management remain in use. + +### Runtime and launch + +The initial adapter targets vLLM 0.15.x V1 with key-first +`[slot, head, Dk, Dv]` FP32 recurrent state. Runtime checks reject other layouts +and versions. Qualification uses source revision `930288170` with Torch +2.9.1/CUDA 12.8 on RTX A6000 GPUs. + +```bash +PYTHONPATH=.:examples/vllm_serve \ +RECIPE_PATH=examples/vllm_serve/linear_attention_state_int8.yaml \ +python examples/vllm_serve/vllm_serve_fakequant.py /path/to/model \ + --tensor-parallel-size 2 --enforce-eager --no-async-scheduling \ + --no-enable-prefix-caching --mamba-cache-dtype float32 +``` + +The example enables signed symmetric dynamic INT8 state QDQ. Set +`num_bits: [4, 3]` in each state quantizer to select FP8 E4M3. Each invocation +receives `[active_sequence, local_head, Dk, Dv]`; `axis: [0, 1]` retains one +scale per sequence and local head, reducing over the whole state matrix. +State tensors remain FP32 after dequantization. + +Dynamic scales need no calibration dataset; `algorithm: null` skips dataset +loading. Weight/activation calibration can use the worker's existing recipe +and calibration loop. Static state calibration is unsupported in this adapter. + +Use `MODELOPT_STATE_PATH=/path/to/modelopt_state.pt` instead of a recipe to +restore quantizer configuration. Reload maps recurrent-state quantizer names +through the model's HF-to-vLLM mapper. Checkpoint weights must already match +the vLLM model architecture. + +### Quantization boundaries + +- Prefill: vLLM gathers initial states and zeros fresh requests. The wrapper + applies `TensorQuantizer` to this tensor and calls the original prefill kernel. +- Decode: the wrapper gathers active native cache slots, quantizes them, writes + them back, and calls the original decode kernel. Other slots are untouched. +- Native kernels compute outputs and update recurrent state normally. No + additional rounding is applied to the kernel's final-state write. +- The next call quantizes that state before reading it. A fresh zero state is + unchanged; prompt-to-decode rounding happens before the first decode call. +- Scheduler-level chunked prefill creates one QDQ boundary per invocation. + Internal kernel chunks do not create extra state-quantization boundaries. + Results can therefore depend on scheduler prompt-chunk sizes. +- Requests, cache-slot reuse, and preemption remain managed by vLLM. The plugin + allocates temporary gathered states, with no additional persistent state cache. +- TP ranks quantize their local heads independently, without scale all-reduce. + +This state-only adapter accepts the default `linear_attention` execution policy +and state quantizer configuration. Enabled W/operand quantizers and nondefault +training, replay, decay, or solve policies are rejected rather than reinterpreted. +The previous experimental replay-serving recipe is replaced by +`linear_attention_state_int8.yaml`; its saved decode policies are unsupported. + +The qualified scope is eager synchronous execution with TP=1/2 and PP=DP=CP=1. +Speculative decoding, prefix caching, state transfer, and CUDA graphs require +separate integration. This is floating-point numerical emulation; model-quality, +capacity, and performance claims require separate measurements. + +### Later prefill GEMM support + +Prefill operand QDQ will be a separate change after an optimized fused kernel +is available. It will reuse the eight numerical sites from the training prefill +implementation and preserve this native-cache wrapper. The PyTorch materialized +backend is not exposed in vLLM by this state-only adapter. State QDQ cadence and +prefill operand QDQ remain independent numerical policies. + ## Known Problems 1. **MCore reload does not use `MODELOPT_STATE_PATH`**; use `QUANT_FILE_PATH` and make sure `QUANT_CFG` matches the quantization recipe used for the original MCore model (otherwise quantizer keys/config won’t align). diff --git a/examples/vllm_serve/fakequant_worker.py b/examples/vllm_serve/fakequant_worker.py index a7d924c06d3..c78abe49985 100644 --- a/examples/vllm_serve/fakequant_worker.py +++ b/examples/vllm_serve/fakequant_worker.py @@ -37,6 +37,7 @@ disable_compilation, post_restore_vllm_parallel_linears, ) +from modelopt.torch.quantization.plugins.vllm_linear_attention import bind_vllm_linear_attention from modelopt.torch.utils import safe_load from modelopt.torch.utils.dataset_utils import get_dataset_dataloader @@ -114,17 +115,22 @@ def _fakequant_run_prolog_worker(self, mlflow_tracker: FakeQuantMlflowTracker) - print("Will load quant, so only do a single sample calibration") quant_config["calib_size"] = 1 - calib_dataloader = get_dataset_dataloader( - dataset_name=quant_config["dataset"], - tokenizer=tokenizer, - batch_size=quant_config["calib_batch_size"], - num_samples=quant_config["calib_size"], - device=self.device, - ) + quant_cfg = get_quant_config(quant_config, model) + calibrate_loop = None + if quant_cfg.get("algorithm", "max") is not None: + calib_dataloader = get_dataset_dataloader( + dataset_name=quant_config["dataset"], + tokenizer=tokenizer, + batch_size=quant_config["calib_batch_size"], + num_samples=quant_config["calib_size"], + device=self.device, + ) + worker_loop = calibrate_fun(calib_dataloader, self) - calibrate_loop = calibrate_fun(calib_dataloader, self) + def calibrate_loop(converted): + bind_vllm_linear_attention(converted, self.model_runner) + worker_loop(converted) - quant_cfg = get_quant_config(quant_config, model) # Before calibration, which is the run this artifact is most wanted for if it dies. mlflow_tracker.log_quant_config(quant_cfg) @@ -134,6 +140,7 @@ def _fakequant_run_prolog_worker(self, mlflow_tracker: FakeQuantMlflowTracker) - quantizer_file_path = quant_config["quant_file_path"] if quantizer_file_path: + bind_vllm_linear_attention(model, self.model_runner) self.model_runner._dummy_run(1) current_state_dict = load_state_dict_from_path(self, quantizer_file_path, model) model.load_state_dict(current_state_dict) @@ -142,6 +149,7 @@ def _fakequant_run_prolog_worker(self, mlflow_tracker: FakeQuantMlflowTracker) - if torch.distributed.is_initialized() and torch.distributed.get_world_size() > 1: torch.distributed.barrier() + bind_vllm_linear_attention(model, self.model_runner) if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0: mtq.print_quant_summary(model) mlflow_tracker.log_quant_summary(model) diff --git a/examples/vllm_serve/linear_attention_state_int8.yaml b/examples/vllm_serve/linear_attention_state_int8.yaml new file mode 100644 index 00000000000..5ca88bfd959 --- /dev/null +++ b/examples/vllm_serve/linear_attention_state_int8.yaml @@ -0,0 +1,12 @@ +metadata: + recipe_type: ptq + description: GDN/KDA INT8 recurrent-state QDQ before native vLLM prefill and decode. +quantize: + quant_cfg: + - quantizer_name: '*' + enable: false + - quantizer_name: '*gdn_state_quantizer' + cfg: {num_bits: 8, type: dynamic, axis: [0, 1], unsigned: false, narrow_range: true} + - quantizer_name: '*kda_state_quantizer' + cfg: {num_bits: 8, type: dynamic, axis: [0, 1], unsigned: false, narrow_range: true} + algorithm: diff --git a/examples/vllm_serve/vllm_reload_utils.py b/examples/vllm_serve/vllm_reload_utils.py index 1bfc7f95fd4..7973d4e71bf 100644 --- a/examples/vllm_serve/vllm_reload_utils.py +++ b/examples/vllm_serve/vllm_reload_utils.py @@ -241,6 +241,19 @@ def _merge_values_require_identical(merged_key: str, key_value_pairs: list[tuple return first_value +def _map_linear_attention_names(values, map_fun): + """Map module paths using weight-path probes without dropping numerical policies.""" + if map_fun is None: + return values + probes = {(name + ".weight" if name else "weight"): value for name, value in values.items()} + mapped = map_fun(probes) + if len(mapped) != len(values) or any( + name != "weight" and not name.endswith(".weight") for name in mapped + ): + raise ValueError("HF-to-vLLM mapping must preserve every linear-attention module") + return {name.removesuffix("weight").removesuffix("."): value for name, value in mapped.items()} + + def convert_dict_to_vllm( state_dict: dict[str, Any], max_or_concat: bool = True, @@ -254,6 +267,15 @@ def convert_dict_to_vllm( max_or_concat: Whether to merge grouped values by taking max/concatenate or require identical map_fun: Function to map the state dict to vLLM format """ + state_dict = dict(state_dict) + linear_quantizers = {} + for key in list(state_dict): + match = re.match(r"^(.*)\.((?:gdn|kda)_(?:state|w)_quantizer(?:\..*)?)$", key) + if match: + module, suffix = match.groups() + mapped = _map_linear_attention_names({module: state_dict.pop(key)}, map_fun) + name, value = next(iter(mapped.items())) + linear_quantizers[f"{name}.{suffix}"] = value # If map_fun is provided, pre-transform quantizer key module-path prefixes so that # HF→vLLM model renames (e.g. backbone.layers → model.layers) are applied before # key grouping (q/k/v → qkv, experts.N.up_proj → experts.w13, etc.). @@ -287,7 +309,7 @@ def convert_dict_to_vllm( _, value = key_value_pairs[0] vllm_state_dict[merged_key] = value if map_fun is None: - return vllm_state_dict + return {**vllm_state_dict, **linear_quantizers} # Quantizer module-path keys (e.g. "layers.0.mlp.gate_proj.input_quantizer") must NOT # go through map_fun (hf_to_vllm_mapper.apply_dict), which maps weight tensor paths and # drops any key it doesn't recognise — including all quantizer keys. Split them out, @@ -295,7 +317,7 @@ def convert_dict_to_vllm( quantizer_keys = {k: v for k, v in vllm_state_dict.items() if "_quantizer" in k} non_quantizer_keys = {k: v for k, v in vllm_state_dict.items() if "_quantizer" not in k} mapped = map_fun(non_quantizer_keys) if non_quantizer_keys else {} - return {**mapped, **quantizer_keys} + return {**mapped, **quantizer_keys, **linear_quantizers} def convert_modelopt_state_to_vllm( @@ -320,6 +342,16 @@ def convert_modelopt_state_to_vllm( modelopt_state_dict = modelopt_state.pop("modelopt_state_dict", []) for idx, current_mode in enumerate(modelopt_state_dict): current_mode_metadata = current_mode[1].pop("metadata", {}) + if "linear_attention" in current_mode_metadata: + current_mode_metadata["linear_attention"] = _map_linear_attention_names( + current_mode_metadata["linear_attention"], map_fun + ) + # Use exact saved module names instead of training-framework selector patterns. + if "quant_cfg" in current_mode[1]["config"]: + current_mode[1]["config"]["linear_attention"] = [ + {"module_name": name, "cfg": policy} + for name, policy in current_mode_metadata["linear_attention"].items() + ] current_mode_quant_state = current_mode_metadata.pop("quantizer_state", {}) if current_mode_quant_state: current_mode_metadata["quantizer_state"] = convert_dict_to_vllm( @@ -367,6 +399,15 @@ def filter_modelopt_state_quantizer_state_for_model( metadata = mode_entry[1].get("metadata", {}) if "quantizer_state" in metadata: saved = metadata["quantizer_state"] + missing_linear = [ + name + for name, state in saved.items() + if re.search(r"(?:gdn|kda)_(?:state|w)_quantizer$", name) + and not state.get("_disabled", False) + and name not in model_keys + ] + if missing_linear: + raise ValueError(f"Saved linear-attention quantizers are missing: {missing_linear}") # Keep keys that exist in the model. Remove disabled quantizers UNLESS they # have registered buffers (e.g. _pre_quant_scale from AWQ/smoothquant on a diff --git a/modelopt/torch/quantization/plugins/__init__.py b/modelopt/torch/quantization/plugins/__init__.py index e1b2ec7a1ae..91b95257cce 100644 --- a/modelopt/torch/quantization/plugins/__init__.py +++ b/modelopt/torch/quantization/plugins/__init__.py @@ -71,6 +71,7 @@ with import_plugin("vllm"): from .vllm import * + from .vllm_linear_attention import * with import_plugin("trl"): from .trl import * diff --git a/modelopt/torch/quantization/plugins/vllm_linear_attention.py b/modelopt/torch/quantization/plugins/vllm_linear_attention.py new file mode 100644 index 00000000000..af98837dd95 --- /dev/null +++ b/modelopt/torch/quantization/plugins/vllm_linear_attention.py @@ -0,0 +1,165 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Recurrent-state QDQ at native vLLM prefill and decode call boundaries.""" + +from functools import partial +from types import FunctionType + +import torch +import vllm +from packaging.version import Version +from vllm.forward_context import get_forward_context + +from ..linear_attention.config import LinearAttentionConfig +from ..nn import QuantModuleRegistry +from .custom import CUSTOM_MODEL_PLUGINS +from .linear_attention import _LinearAttentionQuantMixin + +__all__ = ["bind_vllm_linear_attention"] + + +def bind_vllm_linear_attention(model, model_runner): + """Validate state-only adapters before calibration or generation.""" + layers = [m for m in model.modules() if isinstance(m, _QuantVllmLinearAttention)] + for layer in layers: + layer.validate_linear_attention() + if any(layer.linear_attention_is_enabled for layer in layers): + _validate_runtime(model_runner) + return layers + + +def _validate_runtime(runner): + version = Version(vllm.__version__) + if version.release[:2] != (0, 15): + raise NotImplementedError("Linear-attention fakequant currently requires vLLM 0.15.x V1") + from vllm.model_executor.layers.mamba.mamba_utils import MambaStateShapeCalculator + + shape = MambaStateShapeCalculator.gated_delta_net_state_shape(1, 2, 2, 32, 64, 4)[-1] + if shape[-2:] != (32, 64): + raise NotImplementedError( + "Linear-attention fakequant requires the key-first vLLM state ABI" + ) + config = runner.vllm_config + if not config.model_config.enforce_eager or config.scheduler_config.async_scheduling: + raise ValueError("State fakequant requires --enforce-eager --no-async-scheduling") + if config.speculative_config is not None or config.cache_config.enable_prefix_caching: + raise ValueError("Disable speculative decoding and prefix caching for state fakequant") + if config.cache_config.mamba_cache_mode != "none": + raise ValueError("State fakequant requires mamba_cache_mode='none'") + if config.kv_transfer_config is not None or config.ec_transfer_config is not None: + raise ValueError("State fakequant does not support state transfer") + parallel = config.parallel_config + if any( + getattr(parallel, name, 1) != 1 + for name in ( + "pipeline_parallel_size", + "data_parallel_size", + "prefill_context_parallel_size", + "decode_context_parallel_size", + ) + ): + raise ValueError("State fakequant supports TP with PP=1, DP=1, and CP=1") + + +class _QuantVllmLinearAttention(_LinearAttentionQuantMixin): + def validate_linear_attention(self): + super().validate_linear_attention() + if ( + self._linear_attn_w.is_enabled + or self.linear_attention_config != LinearAttentionConfig() + ): + raise ValueError( + "vLLM supports only state TensorQuantizer at native call boundaries; " + "prefill operand, arithmetic, and decode policies are unsupported" + ) + + def _quantized_state_call(self, native, *args, **kwargs): + state = kwargs.get("initial_state") + if state is not None and self._linear_attn_state.is_enabled: + if state.dtype != torch.float32: + raise ValueError("State fakequant requires an FP32 recurrent cache") + indices = kwargs.get("ssm_state_indices") + if indices is None: + # Native prefill has already initialized fresh requests to zero. + kwargs["initial_state"] = self._linear_attn_state(state) + else: + # Eager, non-speculative decode has one cache slot per active sequence. + indices = indices[: kwargs["cu_seqlens"].numel() - 1].long() + state.index_copy_( + 0, indices, self._linear_attn_state(state.index_select(0, indices)) + ) + return native(*args, **kwargs) + + def _forward_with_state_qdq(self, original, kernel_names, *args, **kwargs): + if not self.linear_attention_is_enabled: + return original(*args, **kwargs) + context = get_forward_context() + if context.attn_metadata is None: + return original(*args, **kwargs) + if context.attn_metadata[self.prefix].spec_sequence_masks is not None: + raise ValueError("Speculative metadata is unsupported for state fakequant") + function = original.__func__ + replacements = { + name: partial(self._quantized_state_call, function.__globals__[name]) + for name in kernel_names + } + # Bind wrappers for this invocation only; every wrapper calls the original kernel. + forward = FunctionType( + function.__code__, + {**function.__globals__, **replacements}, + function.__name__, + function.__defaults__, + function.__closure__, + ) + forward.__kwdefaults__ = function.__kwdefaults__ + return forward(self, *args, **kwargs) + + +class _QuantVllmGDN(_QuantVllmLinearAttention): + def _forward_core(self, *args, **kwargs): + return self._forward_with_state_qdq( + super()._forward_core, + ("chunk_gated_delta_rule", "fused_recurrent_gated_delta_rule"), + *args, + **kwargs, + ) + + +class _QuantVllmKDA(_QuantVllmLinearAttention): + linear_attention_quantizer_names = ("kda_state_quantizer", "kda_w_quantizer") + + def _forward(self, *args, **kwargs): + return self._forward_with_state_qdq( + super()._forward, + ("chunk_kda", "fused_recurrent_kda"), + *args, + **kwargs, + ) + + +def _register_vllm_linear_attention(model): + adapters = { + ("vllm.model_executor.models.qwen3_next", "Qwen3NextGatedDeltaNet"): _QuantVllmGDN, + ("vllm.model_executor.layers.kda", "KimiDeltaAttention"): _QuantVllmKDA, + } + for module in model.modules(): + cls = type(module) + adapter = adapters.get((cls.__module__, cls.__name__)) + if adapter is not None and cls not in QuantModuleRegistry: + QuantModuleRegistry.register({cls: f"vllm_{cls.__name__}"})(adapter) + + +CUSTOM_MODEL_PLUGINS.add(_register_vllm_linear_attention) diff --git a/tests/gpu_vllm/torch/quantization/test_vllm_linear_attention.py b/tests/gpu_vllm/torch/quantization/test_vllm_linear_attention.py new file mode 100644 index 00000000000..76f0d5b1a7c --- /dev/null +++ b/tests/gpu_vllm/torch/quantization/test_vllm_linear_attention.py @@ -0,0 +1,315 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Native vLLM state-boundary QDQ, worker reload, and TP on tiny offline models.""" + +import copy +import gc +import importlib.util +import math +import os +import shutil +from pathlib import Path + +import pytest +import torch +from packaging.version import Version +from transformers import Qwen3NextConfig +from vllm import LLM, SamplingParams +from vllm import __version__ as vllm_version +from vllm.transformers_utils.configs.kimi_linear import KimiLinearConfig + +import modelopt.torch.opt as mto +import modelopt.torch.quantization as mtq +from modelopt.torch.quantization.linear_attention import LinearAttentionConfig +from modelopt.torch.quantization.plugins.vllm_linear_attention import _QuantVllmLinearAttention + +pytestmark = pytest.mark.skipif( + Version(vllm_version).release[:2] != (0, 15), reason="Requires the pinned vLLM 0.15.x ABI" +) +ROOT = Path(__file__).parents[4] + + +def _example_module(name): + spec = importlib.util.spec_from_file_location( + name + "_test", ROOT / "examples/vllm_serve" / (name + ".py") + ) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def _calibrate_worker(worker): + model = worker.model_runner.model + if hasattr(model, "unwrap"): + model = model.unwrap() + layer = next(m for m in model.modules() if isinstance(m, _QuantVllmLinearAttention)) + projection = layer.q_proj if hasattr(layer, "q_proj") else layer.in_proj_qkvz + quantizer = projection.input_quantizer + quantizer.enable() + try: + loop = _example_module("vllm_ptq_utils").calibrate_fun( + [{"input_ids": torch.tensor([[3] * 5, [4] * 5])}], worker + ) + mtq.calibrate(model, algorithm="max", forward_loop=loop) + assert torch.isfinite(quantizer.amax).all() and quantizer.amax.max() > 0 + finally: + quantizer.disable() + return True + + +def _set_state_format_worker(worker, *, fp8): + for layer in worker.model_runner.model.modules(): + if not isinstance(layer, _QuantVllmLinearAttention): + continue + layer._linear_attn_state.num_bits = (4, 3) if fp8 else 8 + layer.validate_linear_attention() + + +def _disable_worker(worker): + for layer in worker.model_runner.model.modules(): + if not isinstance(layer, _QuantVllmLinearAttention): + continue + layer._linear_attn_state.disable() + layer._linear_attn_w.disable() + layer.linear_attention_config = LinearAttentionConfig() + + +def _audit_worker(worker, *, checkpoint=None): + model = worker.model_runner.model + if hasattr(model, "unwrap"): + model = model.unwrap() + layers = [m for m in model.modules() if isinstance(m, _QuantVllmLinearAttention)] + assert layers + if checkpoint is not None and torch.distributed.get_rank() == 0: + torch.save(mto.modelopt_state(model), checkpoint) + reports = [] + for layer in layers: + assert not hasattr(layer, "_linear_attention_cache") + if not hasattr(layer, "_state_audit"): + layer._state_audit = {"prefill": 0, "decode": 0, "changed": 0, "heads": 0} + original = layer._quantized_state_call + + def checked_call(native, *args, _layer=layer, _original=original, **kwargs): + state = kwargs["initial_state"] + indices = kwargs.get("ssm_state_indices") + expected_state = state.clone() + selected = state if indices is None else state.index_select(0, indices.long()) + reference_quantizer = copy.deepcopy(_layer._linear_attn_state) + rounded = reference_quantizer(selected.clone()) + if indices is None: + expected_state = rounded + else: + expected_state.index_copy_(0, indices.long(), rounded) + # Native KDA prefill writes its output into the value-input buffer. + control_args = tuple(x.clone() if isinstance(x, torch.Tensor) else x for x in args) + control_kwargs = { + name: x.clone() if isinstance(x, torch.Tensor) else x + for name, x in kwargs.items() + } + control_kwargs["initial_state"] = expected_state + expected = native(*control_args, **control_kwargs) + native_calls = [] + + def checked_native(*native_args, **native_kwargs): + incoming = native_kwargs["initial_state"] + incoming = ( + incoming if indices is None else incoming.index_select(0, indices.long()) + ) + torch.testing.assert_close(incoming, rounded, atol=0, rtol=0) + native_calls.append(True) + return native(*native_args, **native_kwargs) + + actual = _original(checked_native, *args, **kwargs) + assert native_calls == [True] + for result, control in zip(actual, expected): + if control is not None: + torch.testing.assert_close(result, control, atol=0, rtol=0) + _layer._state_audit["prefill" if indices is None else "decode"] += 1 + _layer._state_audit["changed"] += int(not torch.equal(selected, rounded)) + _layer._state_audit["heads"] = state.shape[1] + return actual + + layer._quantized_state_call = checked_call + reports.append(dict(layer._state_audit)) + return reports + + +def _tiny_model(path, kind): + common = { + "torch_dtype": "bfloat16", + "hidden_size": 128, + "intermediate_size": 256, + "num_attention_heads": 4, + "head_dim": 32, + "vocab_size": 512, + "max_position_embeddings": 128, + } + if kind == "gdn": + config = Qwen3NextConfig( + **common, + num_hidden_layers=2, + num_key_value_heads=2, + linear_num_value_heads=4, + linear_num_key_heads=2, + linear_key_head_dim=32, + linear_value_head_dim=64, + num_experts=4, + num_experts_per_tok=2, + moe_intermediate_size=64, + shared_expert_intermediate_size=64, + mlp_only_layers=[], + full_attention_interval=2, + architectures=["Qwen3NextForCausalLM"], + ) + else: + config = KimiLinearConfig( + **common, + num_hidden_layers=1, + architectures=["KimiLinearForCausalLM"], + linear_attn_config={ + "head_dim": 32, + "num_heads": 4, + "short_conv_kernel_size": 4, + "kda_layers": [1], + "full_attn_layers": [], + }, + ) + config.save_pretrained(path) + shutil.copytree(ROOT / "tests/_test_utils/torch/tokenizer", path, dirs_exist_ok=True) + + +def _generate(llm): + outputs = llm.generate( + [{"prompt_token_ids": [3] * 73}, {"prompt_token_ids": [4]}], + SamplingParams(temperature=0, max_tokens=12, ignore_eos=True, logprobs=1), + use_tqdm=False, + ) + values = [(o.outputs[0].token_ids, o.outputs[0].cumulative_logprob) for o in outputs] + assert all(len(tokens) == 12 and math.isfinite(score) for tokens, score in values) + return values + + +@pytest.mark.parametrize("kind", ["gdn", "kda"]) +@pytest.mark.parametrize("tp", [1, 2]) +@pytest.mark.timeout(360) +def test_fakequant_worker_generation_reference_reload_and_tp(tmp_path, monkeypatch, kind, tp): + if torch.cuda.device_count() < tp: + pytest.skip(f"Requires {tp} GPUs") + monkeypatch.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "1") + monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "spawn") + paths = [str(ROOT), str(ROOT / "examples/vllm_serve"), str(Path(__file__).parent)] + # Spawn restores the parent interpreter search path after reading PYTHONPATH. + for path in reversed(paths): + monkeypatch.syspath_prepend(path) + monkeypatch.setenv("PYTHONPATH", os.pathsep.join([*paths, os.environ.get("PYTHONPATH", "")])) + monkeypatch.setenv( + "RECIPE_PATH", str(ROOT / "examples/vllm_serve/linear_attention_state_int8.yaml") + ) + monkeypatch.delenv("MODELOPT_STATE_PATH", raising=False) + model_path, checkpoint = tmp_path / "model", tmp_path / "state.pt" + _tiny_model(model_path, kind) + kwargs = { + "model": str(model_path), + "load_format": "dummy", + "dtype": "bfloat16", + "max_model_len": 128, + "max_num_seqs": 4, + "max_num_batched_tokens": 64, + "enforce_eager": True, + "enable_prefix_caching": False, + "enable_chunked_prefill": True, + "async_scheduling": False, + "mamba_cache_dtype": "float32", + "kv_cache_memory_bytes": 128 * 1024**2, + "worker_cls": "fakequant_worker.FakeQuantWorker", + "disable_custom_all_reduce": True, + "tensor_parallel_size": tp, + "seed": 17, + } + llm = LLM(**kwargs) + try: + llm.collective_rpc(_audit_worker, kwargs={"checkpoint": str(checkpoint)}) + actual = _generate(llm) + assert _generate(llm) == actual # Fresh requests can reuse the same native state slots. + reports = llm.collective_rpc(_audit_worker) + assert len(reports) == tp + for rank in reports: + assert all( + r["prefill"] > 0 and r["decode"] > 0 and r["changed"] > 0 and r["heads"] == 4 // tp + for r in rank + ) + assert _generate(llm) == actual + llm.collective_rpc(_set_state_format_worker, kwargs={"fp8": True}) + fp8 = _generate(llm) + assert _generate(llm) == fp8 + llm.collective_rpc(_set_state_format_worker, kwargs={"fp8": False}) + assert _generate(llm) == actual + assert all(llm.collective_rpc(_calibrate_worker)) + llm.collective_rpc(_disable_worker) + disabled = _generate(llm) + finally: + llm.llm_engine.engine_core.shutdown() + del llm + gc.collect() + + monkeypatch.delenv("RECIPE_PATH") + monkeypatch.setenv("MODELOPT_STATE_PATH", str(checkpoint)) + llm = LLM(**kwargs) + try: + assert _generate(llm) == actual + finally: + llm.llm_engine.engine_core.shutdown() + del llm + gc.collect() + + monkeypatch.delenv("MODELOPT_STATE_PATH") + llm = LLM(**kwargs) + try: + assert _generate(llm) == disabled + finally: + llm.llm_engine.engine_core.shutdown() + del llm + gc.collect() + + +def test_saved_linear_attention_policy_and_quantizer_names_follow_mapper(): + module = _example_module("vllm_reload_utils") + + metadata = { + "linear_attention": {"backbone.layers.0.attn": {"backend": "fla"}}, + "quantizer_state": {"backbone.layers.0.attn.kda_state_quantizer": {"_disabled": False}}, + } + state = { + "modelopt_state_dict": [("quantize", {"config": {"quant_cfg": []}, "metadata": metadata})] + } + + def mapper(values): + return { + k.replace("backbone.", "model.").replace(".attn.", ".self_attn."): v + for k, v in values.items() + } + + state["modelopt_state_dict"].append( + ("quantize_algo", {"config": {}, "metadata": copy.deepcopy(metadata)}) + ) + result = module.convert_modelopt_state_to_vllm(state, mapper) + converted = result["modelopt_state_dict"][0][1] + assert list(converted["metadata"]["linear_attention"]) == ["model.layers.0.self_attn"] + assert list(converted["metadata"]["quantizer_state"]) == [ + "model.layers.0.self_attn.kda_state_quantizer" + ] + assert converted["config"]["linear_attention"][0]["module_name"] == "model.layers.0.self_attn" + assert "linear_attention" not in result["modelopt_state_dict"][1][1]["config"]