Skip to content

[3/6] Fused Triton GDN/KDA decode QAT - #2562

Draft
kaix-nv wants to merge 6 commits into
kaix/linear-attention-decode-firstfrom
kaix/linear-attention-decode-triton
Draft

kaix-nv wants to merge 6 commits into
kaix/linear-attention-decode-firstfrom
kaix/linear-attention-decode-triton

Conversation

@kaix-nv

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

Copy link
Copy Markdown
Contributor

Linear-attention PR stack — 6 PRs

Order PR Depends on
1/6 #2497 GDN state/W QAT foundation main
2/6 #2519 Torch GDN/KDA decode QAT + INT8 #2497
3/6 #2562 Fused Triton GDN/KDA decode QAT #2519
4/6 #2541 vLLM GDN/KDA state-only fake quantization #2562
5/6 #2503 GDN/KDA prefill GEMM quantization #2541
6/6 #2507 Experimental GDN/KDA approximate inverse #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.

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

import modelopt.torch.quantization as mtq
from modelopt.recipe import load_recipe
from modelopt.torch.quantization.linear_attention import linear_attention_training_phase

recipe = load_recipe("general/ptq/gdn_state_int8_dynamic").quantize
recipe.linear_attention[0].cfg.decode.implementation = "triton"
mtq.quantize(model, recipe)
with linear_attention_training_phase(model, [64]):
    loss = model(input_ids=ids, labels=labels, use_cache=False).loss
    loss.backward()

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 and git diff --check passed. 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.

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

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

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

@codecov

codecov Bot commented Sep 28, 2026

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 1.60000% with 246 lines in your changes missing coverage. Please review.
⚠️ Please upload report for BASE (kaix/linear-attention-decode-first@7b5caf1). Learn more about missing BASE report.

Files with missing lines Patch % Lines
...ch/kernels/quantization/linear_attention/decode.py 0.00% 195 Missing ⚠️
...lopt/torch/quantization/linear_attention/decode.py 5.88% 32 Missing ⚠️
...opt/torch/kernels/quantization/common/fp8_quant.py 0.00% 19 Missing ⚠️
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           
Flag Coverage Δ
unit 57.60% <1.60%> (?)

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.

Signed-off-by: Kai Xu <kaix@nvidia.com>
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-decode-triton branch from 4a182e8 to 7fde632 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-decode-triton branch from 7fde632 to e0bba45 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>
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-decode-triton branch from e0bba45 to 756c330 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>

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