Repository navigation
Conversation
|
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. |
Contributor
|
Important Draft PR not reviewedDraft PRs are not automatically reviewed by default.
To automatically review draft PRs, update your CodeRabbit configuration: reviews:
auto_review:
drafts: true
Comment |
kaix-nv
added this pull request to stack #2510
September 22, 2026 21:27
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## kaix/linear-attention-qat-m2 #2507 +/- ##
================================================================
- Coverage 70.28% 70.28% -0.01%
================================================================
Files 621 622 +1
Lines 68649 68734 +85
================================================================
+ Hits 48253 48308 +55
- Misses 20396 20426 +30
Flags with carried forward coverage won't be shown. Click here to find out more. ☔ View full report in Codecov by Harness. 🚀 New features to boost your workflow:
|
kaix-nv
removed this pull request from stack #2510
September 23, 2026 01:07
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 23, 2026 01:16
34e549e to
ff21e39
Compare
kaix-nv
changed the base branch from
kaix/linear-attention-qat-m3
to
kaix/linear-attention-qat-m2
September 23, 2026 01:16
This was referenced Sep 23, 2026
kaix-nv
added this pull request to stack #2521
September 23, 2026 01:22
kaix-nv
removed this pull request from stack #2521
September 24, 2026 06:17
kaix-nv
added this pull request to stack #2542
September 24, 2026 06:18
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 24, 2026 06:30
ff21e39 to
f3ce747
Compare
kaix-nv
removed this pull request from stack #2542
September 24, 2026 06:31
kaix-nv
added this pull request to stack #2543
September 24, 2026 06:31
Contributor
|
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 24, 2026 18:17
f3ce747 to
99d02c3
Compare
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 25, 2026 01:54
99d02c3 to
9d3c35f
Compare
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
2 times, most recently
from
September 25, 2026 05:11
1dd4bb1 to
999fe60
Compare
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 25, 2026 20:57
999fe60 to
74bd0b9
Compare
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 26, 2026 06:14
74bd0b9 to
b35f426
Compare
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 27, 2026 06:41
b35f426 to
718d3db
Compare
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
2 times, most recently
from
September 28, 2026 06:05
9e4b6ca to
d8c44d9
Compare
kaix-nv
removed this pull request from stack #2543
September 28, 2026 06:06
kaix-nv
added this pull request to stack #2563
September 28, 2026 06:06
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 28, 2026 21:54
d8c44d9 to
dca07e9
Compare
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 29, 2026 00:49
dca07e9 to
e4b37d3
Compare
Signed-off-by: Kai Xu <kaix@nvidia.com>
kaix-nv
force-pushed
the
kaix/linear-attention-qat-m4
branch
from
September 29, 2026 01:29
e4b37d3 to
e6a7d55
Compare
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
removed this pull request from stack #2563
October 5, 2026 06:09
kaix-nv
added this pull request to stack #2563
October 5, 2026 06:11
kaix-nv
removed this pull request from stack #2563
October 5, 2026 06:15
kaix-nv
added this pull request to stack #2658
October 5, 2026 06:15
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Linear-attention series — 7 PRs (1 merged, 6 open)
mainThe 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.
Add an opt-in approximate inverse after the GDN/KDA prefill infrastructure in
#2503. The saved solve policy selects an explicit Neumann polynomial degree and
either Torch or CUDA FP32 execution. Backward differentiates the actual
polynomial; the implementation never silently changes degree or falls back.
Exact triangular solve remains the default.
The current Neumann candidate failed the pinned KDA model-quality screen.
Keep this PR experimental and in draft. The historical study and failed results
are preserved; a successful kernel or optimizer test is not quality recovery.
The solve benchmark shares CUDA timing, warmup, interleaving, and aggregation with the decode/prefill examples. Each benchmark retains its own loss and correctness checks and records the shared helper source hash.
Execution configuration is introduced by #2519, extended for prefill by #2503, and extended here with the solve policy. Numerical test oracles are imported from the test utility package.
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.Usage
Use
{"method": "exact"}(the default) for the supported baseline.See the solve guide.
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:
The tests cover gradients of the actual polynomial, degree validation, residual identities, composed QDQ, save/restore, and packed tails. Those earlier results did not include Megatron/full-FLA-layer integration.
Historical screening of
arcee-ai/AFM-4.5B-Base-KDA-Onlyon fixed WikiText-2 validation data failed the declared NLL margin for degrees 3, 7, 15, and 31. Degree 63 is not a safe fallback, and prior synthetic H100 measurements established no speed advantage. The qualification report retains the original revisions and limits. Those measurements were not rerun; the candidate remains experimental and in draft.Before your PR is "Ready for review"
Additional Information
Depends on #2503 and comes last in the stack. INT8 model-quality comparisons
and Megatron distributed requalification remain pending. #2541 supplies
state-only vLLM integration; serving-time prefill-GEMM quantization is deferred
until an optimized fused kernel is available. A better-conditioned inverse
approximation needs its own numerical and model-quality evidence.