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 #2563
September 28, 2026 06:06
This was referenced Sep 28, 2026
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## kaix/linear-attention-decode-first #2562 +/- ##
=====================================================================
Coverage ? 77.54%
=====================================================================
Files ? 619
Lines ? 68376
Branches ? 0
=====================================================================
Hits ? 53025
Misses ? 15351
Partials ? 0
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:
|
Signed-off-by: Kai Xu <kaix@nvidia.com>
kaix-nv
force-pushed
the
kaix/linear-attention-decode-triton
branch
from
September 28, 2026 21:54
4a182e8 to
7fde632
Compare
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
force-pushed
the
kaix/linear-attention-decode-triton
branch
from
September 29, 2026 00:49
7fde632 to
e0bba45
Compare
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
force-pushed
the
kaix/linear-attention-decode-triton
branch
from
September 29, 2026 01:29
e0bba45 to
756c330
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>
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 PR stack — 6 PRs
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.
What does this PR do?
Type of change: New feature.
Add an optional fused Triton implementation of the GDN/KDA decode-QAT recurrence introduced by #2519. It fuses token/replay updates with FP8 or INT8 state QDQ and supplies a checkpointed custom backward, differentiable continuation carry, and INT8 Hadamard integration. Replay supports encode-once factors.
The kernel uses shared FP8 QDQ primitives, including software E4M3 conversion before SM89. Torch remains the default backend; the KDA example selects Triton. GPU tests compare outputs, reconstructed states, and gradients with the Torch implementation and exercise real KDA layer training/checkpoint restoration.
This is a review split of the existing implementation. The combined tree is identical to the previous #2519 head
30e0659ad478085c79ce2f54b6096d6066b1f814. It does not change the vLLM state-only integration.Usage
The fused implementation supports first-order gradients, key dimensions up to 128, and value blocks 16/32/64/128. Use Torch for higher-order differentiation.
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.
103 passed, 11 skipped on RTX A6000 with Torch 2.9.1+cu128, Triton 3.5.1, fla-core and flash-linear-attention 0.5.1. Ran the complete decode-kernel suite,
TestScaledE4M3, and real KDA training/checkpoint tests. Nine skips require native FP8 conversion on SM89+; two skip unsupported CPU dtypes. Pre-commit andgit diff --checkpassed. Existing tolerances were unchanged.The Torch-only parent independently passed 102 tests, including eight KDA GPU integration tests.
The restacked #2507 tip passed 180 CPU tests. The new Triton tip and all three descendant source trees are byte-for-byte identical to their pre-split counterparts.
Before your PR is "Ready for review"
Additional Information
Depends on #2519. The previously identified low-level Triton replay-anchor gradient issue remains open. This split does not fix it and the change remains a draft. Native FP8 parity on SM89+, distributed training, model-quality recovery, and serving performance require separate qualification.