Skip to content

[5/6] GDN/KDA prefill operand and arithmetic QAT - #2503

Draft
kaix-nv wants to merge 8 commits into
kaix/linear-attention-vllmfrom
kaix/linear-attention-qat-m2
Draft

kaix-nv wants to merge 8 commits into
kaix/linear-attention-vllmfrom
kaix/linear-attention-qat-m2

Conversation

@kaix-nv

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

Copy link
Copy Markdown
Contributor

Linear-attention series — 6 PRs

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

The five open PRs form one native GitHub stack in the order shown. #2497 has landed, so #2519 targets main. #2541 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.

What does this PR do?

Type of change: New feature.

Add configurable GDN and KDA prefill operand fake quantization after the decode
and INT8 state infrastructure in #2519. This PR is stacked on the state-only
vLLM plugin in #2541. Each of the eight logical matmul sites
quantizes its actual transformed operands through ModelOpt TensorQuantizer,
with independent FP8/NVFP4 settings, accumulator rounding schedules, and named
elementwise rounding points. State carry and all rounding sites remain differentiable.
KDA uses causal per-channel gate differences to avoid overflowing inverse-decay factors.

This combines the prefill functionality previously split between this PR and
#2506. The triangular solve is exact; approximate inverse belongs to #2507.
Existing FP8/INT8 state formats, decode policies, explicit phase handoff, and
ModelOpt save/restore remain available. Operand scales and state-write scales
are independent. Working arithmetic remains FP32 inside BF16/FP16 autocast.

GDN and KDA share the dispatch between pure prefill and explicit prefill/decode execution. Operand modules follow ModelOpt temporary-attribute cleanup. The prefill benchmark reuses the decode benchmark measurement helper, with one configuration pass per quantizer.

Execution configuration is introduced by #2519 and extended here for prefill arithmetic. This PR replaces the decode prefix implementation with the shared batched prefill core, removes the superseded prefix helper, and keeps independent numerical oracles under tests.

The QAT entry point is examples/llm_qat/linear_attention/train.py. Usage and numerical contracts live with the example; historical study reports remain beside the scripts in the PRs that introduce them. The training example writes metrics without saving a trained checkpoint, and its source manifest hashes the current implementation files.

This follow-up also introduces the shared decode benchmark and quality-comparison tooling, configs, and comparison tests deferred from #2519. Its study workflow extends the minimal QAT example with fixed-data evaluation and measurement receipts.

Usage

import torch
import modelopt.torch.quantization as mtq

torch.set_float32_matmul_precision("highest")
fp8 = {"num_bits": (4, 3), "type": "dynamic", "axis": (0, 1, 2)}
model = mtq.quantize(model, {
    "quant_cfg": [
        {"quantizer_name": "*", "enable": False},
        {"quantizer_name": "*linear_attn_sites.*", "cfg": fp8},
        {"quantizer_name": "*gdn_w_quantizer", "cfg": fp8},
        {"quantizer_name": "*kda_w_quantizer", "cfg": fp8},
    ],
    "algorithm": None,
    "linear_attention": [{"module_name": "*", "cfg": {"backend": "matmul"}}],
})

The state-read LHS uses the existing GDN/KDA W handle. See
the GDN guide
and the KDA guide
for NVFP4, individual sites, scale domains, arithmetic policies, and framework 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 README-only restack, all runtime code, tests, and study scripts were byte-for-byte identical to its preceding head; the README inherits the state-quantization enablement guidance. Focused README pre-commit hooks, diff checks, and signed-commit verification passed. Model training, distributed integration, and quality measurements were not rerun.

Prior runtime validation on RTX A6000/SM86, Torch 2.9.1+cu128, Triton 3.5.1, and fla-core 0.5.1:

  • 145 CPU tests passed: GDN lifecycle, references, decode/INT8/Hadamard, GDN prefill, and KDA prefill.
  • 117 GPU tests passed: GDN/KDA prefill plus decode, including outputs, final states, gradients, FP8/NVFP4 operands, arithmetic composition, and continuation.
  • Pre-commit, diff checks, and commit-signature verification passed.

The tests preserve independent numerical oracles under the test package and cover checkpoint compatibility, default-disabled handles, packed tails, grouped heads, and autocast. Those earlier results did not include Megatron/full-FLA-layer or model-quality qualification. No serving-speed or quality-recovery claim is made.

Before your PR is "Ready for review"

  • Backward compatible: Yes; new operand handles start disabled and old state formats remain valid.
  • Copied code/dependency guidance: Existing provenance retained; FLA remains an optional pinned dependency.
  • Necessary tests: Yes.
  • Changelog: Updated.
  • Claude approval: Pending; keep draft.

Additional Information

This materialized backend emulates training numerics. It does not provide native
low-precision MMA, compressed states, or a serving speedup. #2541 supplies
state-only vLLM integration; serving-time prefill-GEMM quantization remains
deferred until an optimized fused kernel is available.

Signed-off-by: Kai Xu <kaix@nvidia.com>
@copy-pr-bot

copy-pr-bot Bot commented Sep 22, 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 22, 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.

@codecov

codecov Bot commented Sep 22, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
⚠️ Please upload report for BASE (kaix/linear-attention-vllm@2b9e3a3). Learn more about missing BASE report.

Additional details and impacted files
@@                      Coverage Diff                      @@
##             kaix/linear-attention-vllm    #2503   +/-   ##
=============================================================
  Coverage                              ?   77.50%           
=============================================================
  Files                                 ?      621           
  Lines                                 ?    68539           
  Branches                              ?        0           
=============================================================
  Hits                                  ?    53122           
  Misses                                ?    15417           
  Partials                              ?        0           
Flag Coverage Δ
unit 57.60% <100.00%> (?)

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 removed this pull request from stack #2510 September 23, 2026 01:07
@kaix-nv
kaix-nv added this pull request to stack #2520 September 23, 2026 01:07
@kaix-nv
kaix-nv removed this pull request from stack #2520 September 23, 2026 01:08
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-m2 branch from 86c4600 to bd1f350 Compare September 23, 2026 01:16
@kaix-nv kaix-nv changed the title [QAT] Add GDN prefill operand and arithmetic emulation [3/4] GDN/KDA prefill operand and arithmetic QAT Sep 23, 2026
@kaix-nv
kaix-nv changed the base branch from kaix/linear-attention-qat-m1 to kaix/linear-attention-decode-first September 23, 2026 01:16
@kaix-nv
kaix-nv added this pull request to stack #2521 September 23, 2026 01:22
@kaix-nv kaix-nv changed the title [3/4] GDN/KDA prefill operand and arithmetic QAT [4/5] GDN/KDA prefill operand and arithmetic QAT Sep 24, 2026
@kaix-nv
kaix-nv removed this pull request from stack #2521 September 24, 2026 06:17
@kaix-nv
kaix-nv added this pull request to stack #2542 September 24, 2026 06:18
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-m2 branch from bd1f350 to 580c3dc Compare September 24, 2026 06:30
@kaix-nv
kaix-nv removed this pull request from stack #2542 September 24, 2026 06:31
@kaix-nv
kaix-nv changed the base branch from kaix/linear-attention-decode-first to kaix/linear-attention-vllm September 24, 2026 06:31
@kaix-nv
kaix-nv added this pull request to stack #2543 September 24, 2026 06:31
@github-actions

github-actions Bot commented Sep 24, 2026 •

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-2503/

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

@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-m2 branch 2 times, most recently from 207e918 to 2dd8c1d Compare September 25, 2026 01:54
@kaix-nv
kaix-nv removed this pull request from stack #2543 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 [4/5] GDN/KDA prefill operand and arithmetic QAT [5/6] GDN/KDA prefill operand and arithmetic QAT Sep 28, 2026
Signed-off-by: Kai Xu <kaix@nvidia.com>
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-m2 branch from 86fa9e1 to 29687f3 Compare September 28, 2026 21:54
Warm standalone and Megatron forward/backward kernels before functional tests. Remove the 300-second overrides so test calls retain the default 120-second limit and report execution separately from compilation.

Signed-off-by: Kai Xu <kaix@nvidia.com>
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-m2 branch from 29687f3 to e893646 Compare September 29, 2026 00:49
Signed-off-by: Kai Xu <kaix@nvidia.com>
Signed-off-by: Kai Xu <kaix@nvidia.com>
Signed-off-by: Kai Xu <kaix@nvidia.com>
Signed-off-by: Kai Xu <kaix@nvidia.com>
Signed-off-by: Kai Xu <kaix@nvidia.com>
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-m2 branch from e893646 to bd2be2b Compare September 29, 2026 01:29
**kwargs,
)
recurrent = recurrent_delta_rule_reference(*args, **kwargs)
chunk = chunk_kda_reference(*args, **kwargs)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Does this actually test against the triton kernel? or is the reference a pure pytorch implementation?

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 [5/6] GDN/KDA prefill operand and arithmetic QAT [6/7] GDN/KDA prefill operand and arithmetic QAT Oct 5, 2026
@kaix-nv
kaix-nv removed this pull request from stack #2658 October 8, 2026 18:41
@kaix-nv
kaix-nv added this pull request to stack #2714 October 8, 2026 18:41
@kaix-nv kaix-nv changed the title [6/7] GDN/KDA prefill operand and arithmetic QAT [5/6] GDN/KDA prefill operand and arithmetic QAT Oct 8, 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.

2 participants