Conversation
Signed-off-by: Kai Xu <kaix@nvidia.com>
|
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: trueComment |
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-m4 #2509 +/- ##
================================================================
- Coverage 70.52% 70.38% -0.14%
================================================================
Files 617 620 +3
Lines 68119 68559 +440
================================================================
+ Hits 48039 48257 +218
- Misses 20080 20302 +222
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
Contributor
|
This was referenced Sep 23, 2026
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>
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.
Superseded by #2519: Decode is delivered before prefill and inverse, with signed INT8 state and replay-anchor support. Prefill composition follows in #2503 and inverse in #2507. This closed PR and its branch retain the original qualification history.
What does this PR do?
Type of change: New feature, new example, new tests, documentation.
Adds experimental differentiable GDN/KDA decode policies for training: per-token FP8 state writes, optional log-retention grid rounding, and anchor-plus-encoded-update replay. Workloads supply explicit prefix lengths, preserving gradients across prefill/decode handoff, anchor refresh, and continuation. Policies persist through ModelOpt save/restore.
A Torch reference defines write/readout order and codec metadata. A fused FP32 Triton forward/backward implements token and encode-once replay with internal state checkpoints. A fixed key-reduction tree and matching dynamic FP8 scale arithmetic prevent tiny arithmetic differences from flipping rounding ties and accumulating along near-unit-decay trajectories. Replay factors are computed from the current reconstructed trajectory, never a teacher state.
The default execution path is unchanged. Serving caches, compressed storage, native low-precision speedups, and higher-order fused differentiation are outside this draft. The study keeps the exact solve; it does not enable the Neumann candidate rejected in #2507.
Usage
Use
*gdn_state_quantizerfor GDN. Token mode omits the replay settings. Seedocs/linear_attention_decode.mdfor carry/readout contracts and unsupported cases.Testing
1e-4.Pinned Arcee KDA / Wikitext-2 pilot: 32 validation blocks selected log grid 1/256 before test access. Exact, token-state FP8, state FP8 + selected grid, and replay each completed 32 attention-only updates with finite gradients for all trainable parameters and verified weight changes. On 64 held-out suffix-scored blocks, post-training perplexities were 12.6413 / 12.6567 / 12.6530 / 12.6466. All approximations met the 0.02 NLL pilot margin; the largest 95% upper bound was 0.002703. This is a short-context, single-model pilot, not broad quality recovery.
For complete
[1,257,4,128]prefix/suffix training on H100 (prefix 64, three warmups, 20 interleaved samples), KDA token/decay/replay took 11.265 / 11.310 / 11.569 ms versus 228.732 / 235.972 / 327.361 ms for matching Torch references. Exact FLA with BF16 Q/K/V took 2.347 ms. Numerical outputs/state matched the references and maximum gradient relative-norm error was below 3.4e-7. GDN results and paired intervals are indocs/linear_attention_decode_study.md; raw receipts bind retained source snapshots. Final import/comment cleanup preserves the non-import AST and was regression-tested; the benchmark harness differs only in typing import-name ordering.Additional tests cover full-width state/value tails at codec blocks 32/128 and composition of prefix FP8, Neumann solve, elementwise rounding, and replay across packed/empty sequences and state layouts. These synthetic composition tests do not promote the Neumann algorithm for model quality.
Before your PR is "Ready for review"
CONTRIBUTING.md: N/A No new dependency or copied third-party implementation.Additional Information
Stacked on #2507 (M4), following #2506 (KDA prefill), #2503 (GDN materialized prefill), and #2497 (GDN state/W QDQ). This draft models numerical behavior for QAT and has no inference performance or compressed-storage claim.