Skip to content

[5/7] vLLM GDN/KDA state-only fake quantization - #2541

Draft
kaix-nv wants to merge 1 commit into
kaix/linear-attention-decode-tritonfrom
kaix/linear-attention-vllm
Draft

kaix-nv wants to merge 1 commit into
kaix/linear-attention-decode-tritonfrom
kaix/linear-attention-vllm

Conversation

@kaix-nv

@kaix-nv kaix-nv commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

Linear-attention series — 7 PRs (1 merged, 6 open)

Order PR Depends on
1/7 #2497 GDN state/W QAT foundation main
2/7 #2519 Torch GDN/KDA decode QAT + INT8 #2497
3/7 #2657 Megatron Bridge linear attention QAT/QAD example #2519
4/7 #2562 Fused Triton GDN/KDA decode QAT #2657
5/7 #2541 vLLM GDN/KDA state-only fake quantization #2562
6/7 #2503 GDN/KDA prefill GEMM quantization #2541
7/7 #2507 Experimental GDN/KDA approximate inverse #2503

The six open PRs form native GitHub stack #2658 in the order shown; #2497 is retained as the merged foundation in this seven-PR series. #2497 is merged, so #2519 targets main. #2657 contains the training example split from #2519; #2562 now targets #2657. Rebase each remaining descendant after its immediate parent merges.

#2541 applies TensorQuantizer before native vLLM prefill/decode calls. Serving-time prefill-GEMM quantization remains deferred until an optimized fused kernel is available. #2506 and #2509 are superseded and closed.

What does this PR do?

Type of change: New feature.

Stacked on #2562 (kaix/linear-attention-decode-triton). Add a state-only ModelOpt fake-quant plugin for vLLM GDN and KDA. It applies TensorQuantizer to the incoming recurrent state immediately before each native prefill or decode call. Native kernels and cache management remain in use; this PR changes no CUDA or Triton kernels.

Prefill quantizes the initialized state tensor passed to the original chunk kernel. Decode gathers only active cache slots, applies QDQ, and writes them back before the original recurrent kernel. Each slot/head has its own dynamic scale over [Dk,Dv]. FP8 E4M3 and signed symmetric INT8 are supported, with FP32 dequantized state.

This replaces the earlier draft's custom attention execution and separate request cache. There is no additional persistent state allocation or worker memory reservation. Quantization happens once per native invocation, including scheduler-level prompt continuations; it does not round every internal prefill chunk or the final-state write. Quantizer configuration and state names survive save/restore through the HF-to-vLLM mapper.

The worker validates adapter policy before calibration or warmup and after loading quantizer state. It discovers adapters from the model instead of retaining a second list; shared runtime capability checks run once per binding.

The execution-config definition comes from #2519. This serving adapter continues to accept only the default execution policy and state quantizers; it does not enable the training prefill-GEMM or replay paths. The native integration test now propagates its import paths to spawned workers as well as through PYTHONPATH.

The serving example README includes the state-quantization launch command, numerical boundaries, and runtime requirements.

Prefill/decode quantization boundary

State QDQ applies to both prefill and decode, specifically at native vLLM invocation boundaries:

  • Prefill: quantize the incoming initialized recurrent state once before each native prefill invocation. With scheduler-level chunked prefill, each invocation creates one QDQ boundary; internal kernel chunks do not.
  • Decode: gather and quantize only the active recurrent-cache slots immediately before each native decode invocation. Inactive slots are unchanged.
  • State write: native kernels update and write recurrent state in FP32 without additional final-write quantization. The next prefill or decode invocation quantizes that state before reading it.

This is floating-point fake QDQ for numerical qualification. It does not store the persistent recurrent cache in packed INT8 or FP8, execute state updates with low-precision MMA, or establish memory-capacity, bandwidth, or latency savings.

Usage

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 recipe sets algorithm: null: dynamic state scales require no calibration dataset. Set state num_bits: [4, 3] for FP8. Existing worker weight/activation calibration remains available. Use MODELOPT_STATE_PATH instead of RECIPE_PATH to restore saved quantizer configuration.

The state-only adapter rejects nondefault execution policies, including the previous draft's replay policy. Use the new state-only recipe. See examples/vllm_serve/README.md#linear-attention-state-quantization for the exact rounding cadence and runtime limits.

Testing

Compilation-fixture update: this branch is restacked on #2497's separate follow-up commit 95766de709ac. The changed test modules passed at #2497 (26 passed, 20 hardware skips in each cold/warm run) and at the #2507 stack tip (64 passed, 20 hardware skips on two RTX A6000 GPUs); intermediate PRs were not separately rerun. Functional calls created no new tracked kernel binaries and kept the default 120-second cap. All six source trees match the validated trees, and commit hooks passed. Native FP8-state/Hopper cases remain hardware-gated.

Earlier scope-specific validation follows.

For the earlier documentation/example amendment, pre-commit, Markdown links and anchors, command/Python syntax, source-manifest readability, and stale-path checks passed. Runtime kernels were not changed; model training, distributed integration, and quality studies were not rerun for this amendment.

Prior runtime validation on two RTX A6000 GPUs, Torch 2.9.1+cu128, and Triton 3.5.1:

  • 13 CPU export/reload tests passed.
  • All 5 native vLLM integration tests passed (578.19 seconds): GDN and KDA at TP=1/2 plus saved-state name mapping. Tests exercise INT8 and FP8 state QDQ, native-kernel controls, inactive cache slots, prompt continuations, repeated generation, calibration, checkpoint reload, and disabled-path agreement.
  • Pre-commit, diff checks, and commit-signature verification passed.

The integration run used the clean pinned vLLM checkout 930288170c31e8568290fff407dca8caf17d16ad, whose recurrent-state ABI is key-first, with existing local compiled artifacts. The initially selected editable vLLM checkout used value-first state and correctly failed the runtime guard; that run is not included in the passing result.

The passing run used the existing local pytest harness to set gpu_memory_utilization=0.04, with the same 128 MiB cache budget and unchanged numerical assertions. This avoids requiring nearly all GPU memory for tiny synthetic models. NCCL_P2P_DISABLE=1 is required on this host. These are validation settings, not changes to the serving worker or kernels.

PYTHONPATH=. python -m pytest -q \
  tests/unit/torch/export/test_vllm_fakequant_hf.py \
  tests/unit/torch/export/test_vllm_quantizer_reload.py

NCCL_P2P_DISABLE=1 \
PYTHONPATH=.:tmp/vllm-runtime:examples/vllm_serve:tmp/design-review \
python -m pytest -p low_memory_vllm -q \
  tests/gpu_vllm/torch/quantization/test_vllm_linear_attention.py

The local low_memory_vllm fixture wraps this test module's LLM constructor with functools.partial(LLM, gpu_memory_utilization=0.04). It is not part of the PR. Models use tiny offline Qwen3-Next/Kimi Linear configurations and synthetic weights; this is functional qualification, with no pretrained-quality or performance claim.

Before your PR is "Ready for review"

  • Is this change backward compatible?: Opt-in; disabled-path agreement is tested. The unmerged draft's replay-serving recipe and execution policies are intentionally replaced by state-only QDQ.
  • If you copied code or added a dependency, did you follow CONTRIBUTING.md?: N/A; no copied third-party kernel or new dependency. Wrappers call native vLLM functions.
  • Did you write necessary tests?: Yes; native-kernel and worker integration coverage.
  • Did you update Changelog?: Yes.
  • Did you get Claude approval?: No; keep draft.

Additional Information

The supported runtime is vLLM 0.15.x with key-first FP32 state, eager synchronous execution, TP=1/2 and PP=DP=CP=1. Speculative decoding, prefix caching, state transfer, and CUDA graphs remain unsupported.

A separate prefill-GEMM PR is deferred until an optimized fused kernel is available. It will reuse #2503's eight numerical sites; this adapter does not expose the materialized PyTorch prefill backend. Megatron training/backward qualification and vLLM forward/cache qualification remain separate.

@copy-pr-bot

copy-pr-bot Bot commented Sep 24, 2026

Copy link
Copy Markdown

Auto-sync is disabled for draft pull requests in this repository. Workflows must be run manually.

Contributors can view more details about this message here.

@coderabbitai

coderabbitai Bot commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

Important

Draft PR not reviewed

Draft PRs are not automatically reviewed by default.

  • Trigger a manual review

To automatically review draft PRs, update your CodeRabbit configuration:

reviews:
  auto_review:
    drafts: true
  • Autopilot · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

Comment @coderabbitai help to get the list of available commands.

@kaix-nv kaix-nv changed the title Add vLLM GDN/KDA decode fake quantization Add vLLM GDN/KDA state-only fake quantization Sep 24, 2026
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from ca9f302 to 087123e Compare September 24, 2026 06:06
@kaix-nv kaix-nv changed the title Add vLLM GDN/KDA state-only fake quantization [3/5] vLLM GDN/KDA state-only fake quantization Sep 24, 2026
@kaix-nv
kaix-nv added this pull request to stack #2543 September 24, 2026 06:31
@codecov

codecov Bot commented Sep 24, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 0% with 73 lines in your changes missing coverage. Please review.
✅ Project coverage is 77.46%. Comparing base (4a182e8) to head (2b9e3a3).

Files with missing lines Patch % Lines
...orch/quantization/plugins/vllm_linear_attention.py 0.00% 72 Missing ⚠️
modelopt/torch/quantization/plugins/__init__.py 0.00% 1 Missing ⚠️
Additional details and impacted files
@@                           Coverage Diff                           @@
##           kaix/linear-attention-decode-triton    #2541      +/-   ##
=======================================================================
- Coverage                                77.54%   77.46%   -0.09%     
=======================================================================
  Files                                      619      620       +1     
  Lines                                    68376    68449      +73     
=======================================================================
  Hits                                     53025    53025              
- Misses                                   15351    15424      +73     
Flag Coverage Δ
unit 57.53% <0.00%> (-0.07%) ⬇️

Flags with carried forward coverage won't be shown. Click here to find out more.

☔ View full report in Codecov by Harness.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from 087123e to f9d3b35 Compare September 24, 2026 18:17
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from f9d3b35 to 9ccae1f Compare September 25, 2026 01:54
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from 9ccae1f to 3c932d0 Compare September 25, 2026 04:26
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from 3c932d0 to 830a114 Compare September 25, 2026 05:11
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from 830a114 to 1133ecf Compare September 25, 2026 20:57
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch 2 times, most recently from cf7c163 to 2989a37 Compare September 27, 2026 06:41
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from 2989a37 to e22ded5 Compare September 28, 2026 03:29
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from e22ded5 to 2b9e3a3 Compare September 28, 2026 06:05
@kaix-nv
kaix-nv removed this pull request from stack #2543 September 28, 2026 06:06
@kaix-nv
kaix-nv changed the base branch from kaix/linear-attention-decode-first to kaix/linear-attention-decode-triton September 28, 2026 06:06
@kaix-nv
kaix-nv added this pull request to stack #2563 September 28, 2026 06:06
@kaix-nv kaix-nv changed the title [3/5] vLLM GDN/KDA state-only fake quantization [4/6] vLLM GDN/KDA state-only fake quantization Sep 28, 2026
@github-actions

Copy link
Copy Markdown
Contributor
PR Preview Action v1.8.1

QR code for preview link

🚀 View preview at
https://NVIDIA.github.io/Model-Optimizer/pr-preview/pr-2541/

Built to branch gh-pages at 2026-09-28 06:15 UTC.
Preview will be ready when the GitHub Pages deployment is complete.

@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from 2b9e3a3 to 6daaeea Compare September 28, 2026 21:54
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from 6daaeea to 141688f Compare September 29, 2026 00:49
Signed-off-by: Kai Xu <kaix@nvidia.com>
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-vllm branch from 141688f to bae8fa2 Compare September 29, 2026 01:29
kaix-nv added a commit that referenced this pull request Oct 2, 2026
<!-- linear-attention-stack:start -->
**Linear-attention PR stack — 6 PRs**

| Order | PR | Depends on |
| --- | --- | --- |
| 1/6 | [#2497 GDN state/W QAT
foundation](#2497) | main
|
| 2/6 | [#2519 Torch GDN/KDA decode QAT +
INT8](#2519) | #2497 |
| 3/6 | [#2562 Fused Triton GDN/KDA decode
QAT](#2562) | #2519 |
| 4/6 | [#2541 vLLM GDN/KDA state-only fake
quantization](#2541) |
#2562 |
| 5/6 | [#2503 GDN/KDA prefill GEMM
quantization](#2503) |
#2541 |
| 6/6 | [#2507 Experimental GDN/KDA approximate
inverse](#2507) | #2503 |

All six PRs form native GitHub stack #2563 in the order shown above.
#2541 applies TensorQuantizer before native vLLM prefill/decode calls. A
separate vLLM prefill-GEMM PR waits for an optimized fused kernel. #2506
and #2509 are superseded and closed.
<!-- linear-attention-stack:end -->

### What does this PR do?

Type of change: new feature

GatedDeltaNet training keeps recurrent states inside a chunked kernel,
so projection quantizers cannot emulate rounding at state boundaries.
This PR adds dynamic per-tile FP8 E4M3 fake QDQ to the recurrent state
and independent dynamic FP8 fake QDQ to WY-transformed W activations,
with identity straight-through gradients for QAT/QAD.

Both sites use the standard `quant_cfg` interface and start disabled.
State QDQ uses 64-token chunks and recomputes `amax` at each boundary
over each full-key by 64-value-column tile, independently per sequence
and head. Each tile has its own scalar scale (`amax / 448`, with a zero
guard); `fp8_scalar_qdq` applies that supplied scale rather than
choosing tensor-wide grouping. W grouping is applied by
`TensorQuantizer`. Quantizer settings use normal ModelOpt checkpoint
state. There is no `QuantizeConfig.linear_attention` field in this PR;
#2519 introduces execution policies for decode and ReplaySSM, and later
PRs extend them for prefill and approximate inverse. Configurations or
checkpoints from earlier experimental drafts that use those execution
policies require #2519; those selecting Triton decode also require
#2562.

The Megatron adapter supports the direct-forward and older split-forward
call layouts, restores the original kernel when disabled, and removes
temporary quantizer attributes on export. Independent recurrent/chunk
numerical references live under `tests/_test_utils/torch/quantization/`;
shared runtime capability checks live in `linear_attention/utils.py`.

The fused path requires `fla-core==0.5.1` and chunk size 64. State FP8
emulation requires SM89 or newer. The Hopper path has additional
dtype/TileLang restrictions enforced before launch. This PR simulates
numerical error; it does not add compressed state storage or faster
inference.

### Usage

```python
import modelopt.torch.quantization as mtq

model = mtq.quantize(model, {
    "quant_cfg": [
        {"quantizer_name": "*", "enable": False},
        {"quantizer_name": "*gdn_state_quantizer",
         "cfg": {"num_bits": (4, 3), "type": "dynamic", "axis": (0, 1)}},
        {"quantizer_name": "*gdn_w_quantizer",
         "cfg": {"num_bits": (4, 3), "type": "dynamic", "axis": (0, 1, 2)}},
    ],
    "algorithm": None,
})
# Continue with the framework's normal forward/backward/optimizer steps.
```

Dynamic scales require no calibration.

### Testing

The focused GPU suite contains four cases: three BF16 numerical
forward/backward checks (disabled, W QDQ, and state+W QDQ) using one
shared shape, plus one single-rank, one-layer Megatron QAT/checkpoint
test. The Megatron test checks quantizer enable/disable behavior,
checkpoint restore, gradients, and an optimizer update; it enables state
QDQ when the GPU supports native FP8 conversion. Compilation runs in
setup fixtures, and functional calls retain the normal 120-second
timeout. There are no dtype, layout, tile-width, or parallelism sweeps.

The pinned FLA/TileLang/TVM-FFI dependencies live in the `dev-fla`
optional extra, installed by both GPU nox sessions.

Validation of the consolidated changes on RTX A6000 (SM86), Python
3.12.8, Torch 2.9.1+cu128, Triton 3.5.1, fla-core 0.5.1, TileLang 0.1.8,
Megatron Core 0.19.2, and Transformer Engine 2.16.0:

- Cold and warm focused runs: **3 passed, 1 hardware skip** each. The
state+W numerical case requires SM89+; the local Megatron test exercised
W QDQ.
- Fresh Triton/TileLang cache: **363.09s total**, including setup and
teardown. Kernel setup took 66.38s + 44.46s; Megatron setup, including
shared extension setup and worker startup, took 245.26s. Functional
calls totaled about 2.56s.
- Same cache, new pytest process: **38.20s total**, with about **2.41s
in functional calls**.
- Pre-commit checks passed for the four changed files. Dependency-group
wiring and installed pinned versions were checked.

```bash
PYTHONPATH=. python -m pytest -q \
  tests/gpu/torch/kernels/quantization/linear_attention/test_fla_chunk_gated_delta_rule.py \
  tests/gpu_megatron/torch/quantization/plugins/test_megatron_gated_delta_net.py \
  --durations=0
```

These timings describe local test setup and execution, not inference
performance. Native FP8 state QDQ and Hopper still require suitable
GPU/CI runs. This minimal suite does not qualify tensor/context/pipeline
parallelism, checkpoint resharding, or model-quality recovery. Mamba
compilation coverage is tracked separately in #2572.

### Before your PR is "*Ready for review*"

Contributor and security guidance reviewed. Commits are signed and
signed off.

- Is this change backward compatible?: ✅ Disabled-by-default quantizers,
standard-recipe exclusions, and legacy-checkpoint coverage; enabled
experimental configurations have explicit capability restrictions.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ❌ Internal
third-party approval tracking still needs confirmation. Upstream
attribution, MIT/Apache headers, `LICENSE` notice, and license-hook
exclusions are included. FLA/TileLang and TVM-FFI license files were
reviewed.
- Did you write any new necessary tests?: ✅ Numerical, gradient,
conversion/checkpoint, and real framework tests.
- Did you update Changelog?: ✅ Experimental quantization feature entry.
- Did you get Claude approval on this PR?: ❌ Bot feedback addressed or
discussed; renewed approval pending.

### Additional Information

Related: #2455. This is the first integration slice and does not assume
#2455 has merged. Later milestones will extend the numerical boundaries
after choosing their approximation contracts.


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Added experimental dynamic FP8 fake quantization for GatedDeltaNet
recurrent states and WY activations during training.
* Added PTQ configuration options for state and WY activation
quantization. State quantization requires an SM89-or-newer GPU; the
fused path requires `fla-core==0.5.1` and a chunk size of 64.
* **Bug Fixes**
* Improved quantizer configuration validation and restoration for
linear-attention models.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Kai Xu <kaix@nvidia.com>
@kaix-nv
kaix-nv removed this pull request from stack #2563 October 5, 2026 06:09
@kaix-nv
kaix-nv added this pull request to stack #2563 October 5, 2026 06:11
@kaix-nv
kaix-nv removed this pull request from stack #2563 October 5, 2026 06:15
@kaix-nv
kaix-nv added this pull request to stack #2658 October 5, 2026 06:15
@kaix-nv kaix-nv changed the title [4/6] vLLM GDN/KDA state-only fake quantization [5/7] vLLM GDN/KDA state-only fake quantization Oct 5, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant