Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions CHANGELOG.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
6 changes: 6 additions & 0 deletions examples/llm_qat/linear_attention/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
77 changes: 77 additions & 0 deletions examples/vllm_serve/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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).
Expand Down
26 changes: 17 additions & 9 deletions examples/vllm_serve/fakequant_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)

Expand All @@ -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)
Expand All @@ -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)
Expand Down
12 changes: 12 additions & 0 deletions examples/vllm_serve/linear_attention_state_int8.yaml
Original file line number Diff line number Diff line change
@@ -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:
45 changes: 43 additions & 2 deletions examples/vllm_serve/vllm_reload_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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.).
Expand Down Expand Up @@ -287,15 +309,15 @@ 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,
# apply map_fun only to non-quantizer keys, then merge back.
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(
Expand All @@ -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(
Expand Down Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions modelopt/torch/quantization/plugins/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,7 @@

with import_plugin("vllm"):
from .vllm import *
from .vllm_linear_attention import *

with import_plugin("trl"):
from .trl import *
Expand Down
Loading