Skip to content

[2/6] Torch GDN/KDA decode QAT with INT8 recurrent state - #2519

Open
kaix-nv wants to merge 22 commits into
mainfrom
kaix/linear-attention-decode-first
Open

kaix-nv wants to merge 22 commits into
mainfrom
kaix/linear-attention-decode-first

Conversation

@kaix-nv

@kaix-nv kaix-nv commented Sep 23, 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 recurrent-state QAT for Megatron-Core GDN and KDA. Each sequence uses a chunked prefill prefix followed by a recurrent suffix. TensorQuantizer applies fake QDQ at the handoff and subsequent state writes, so training includes the recurrent quantization effects encountered during decode.

An explicit linear-attention policy selects the native forward arithmetic and state/checkpoint schedule. Training uses native forward values with differentiable FLA/Torch surrogate gradients and identity STE through QDQ; this does not claim an exact backward for inference kernels.

  • precision="vllm" selects the supported installed Triton APIs. GDN qualification covers pure single-token decode with matching Triton prefill and packed-decode settings. Generic KDA selects standalone FLA/Triton APIs without model-level serving parity.
  • precision="vllm_kimi_k3" selects the Kimi-K3/Kimi-Linear Triton profile. Serving must use Triton prefill/decode, Q/K normalization and gate_lower_bound=None.
  • precision="replayssm" requires an optional private quantized-ReplaySSM fork, unavailable in public vLLM releases. Without access to it, use a block32 recipe. The fork profile uses Hadamard checkpoints with replay_window=1 by default.

The public training phase supplies per-sequence prefix and valid-token lengths, including packed padding. Active state QAT requires native Megatron _compute_gates and _forward_compute hooks for raw gates and packed sequence lengths; Core 0.19.2 lacks them and is rejected before forward execution. Selective GDN recompute captures the phase for backward on supported Core builds. Active state QAT rejects context parallelism and full-layer recompute. Disabled state quantizers outside a phase preserve the original module path.

This change supports floating-storage state fake quantization. It does not add a vLLM server or native compressed cache. #2541 owns serving integration; #2657 owns the training example. Prefill GEMM quantization remains a follow-up. Descendant rebases remain deferred until their immediate parent merges.

Usage

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

cfg = load_recipe("general/ptq/linear_attention_state_int8_block32_dynamic").quantize
model = mtq.quantize(model, cfg)
with linear_attention_training_phase(model, prefill_lengths=[64], sequence_lengths=[128]):
    loss = training_step(model, batch)
loss.backward()

This expects an initialized Megatron model and a loss on the decode suffix. Training needs compatible vLLM and FLA packages. Validation/calibration loops must also establish the appropriate phase. A full-prefill pass does not measure decode state-QDQ quality. Compare against a converted model with the same serving policy and disabled state quantizers.

Testing

Current head: 03f6076dd8. CI results for this head are pending. The prior GPU CI run on 140735b962 had these results:

  • vLLM 0.20 on sm_120: CUDA misaligned access during GDN setup, followed by failures from the poisoned context.
  • vLLM 0.30 on sm_120: linear-attention tests passed, but the whole job exceeded its 15-minute limit.
  • NeMo 26.08: the earlier GDN test passed, but that does not establish serving parity on Core builds lacking the gate/packed-forward hooks. This update rejects that unsupported execution path. NAS test_mamba_search_space timed out; KDA skips because the image lacks its module.
  • Regression and GPT-OSS checks passed on the prior head.

Local validation of the code published in 03f6076dd8, on RTX A6000:

  • Focused CPU policy/checkpoint/state-QDQ and recipe documentation tests: 44 passed, including missing disabled handles with policy metadata, unconverted GDN/KDA recipe rejection, plain-KDA gate rejection and the runtime hardware guard.
  • vLLM 0.30 / Torch 2.13 / Triton 3.7.1: 11 passed, 3 skipped. Native bit-exact output/state comparisons, handoff gradients, mismatch rejection, K=192 layout validation, GDN/KDA padding gradients and an independent prefill-only reference pass. Only the private native ReplaySSM cases skip. This used a warm local kernel cache; it does not establish CI compilation time.
  • Megatron-Core 0.19.2: 1 passed, 1 skipped. Checkpoint restore and disabled baseline remain usable; active serving QAT raises before the native forward, including the alternate entry, and the full QAT test skips. Missing raw gates are rejected at kernel dispatch.
  • Megatron source ac100f773f9d / Torch 2.9.1 / Transformer Engine 2.16 / vLLM 0.15.2.dev: 2 GDN tests passed, including QAT/sharded restore, packed recompute and the raw-gate/alternate-entry guards.
  • Megatron KDA: 1 passed on the same Core source with vLLM 0.30 / Torch 2.13 / Transformer Engine 2.16.1 / cuBLAS 13.2.2.2. The real layer test covers state QAT, calibration, microbatch-specific recompute gradients, eval gradients, disabled baseline and an optimizer update. Cold setup (including shared CUDA extensions) took 754 seconds; the test call took 0.75 seconds. Earlier dependency-loading failures are preserved in the local logs. This is local adapter coverage; the NeMo 26.08 CI image still lacks KDA.
  • Scoped pre-commit and git diff --check: passed.

The runtime rejects sm_120+ with Triton < 3.7 before state-QAT kernels launch. The whole training test module also skips this combination, including layout probes. The observed misaligned-address fault has not been isolated to a specific kernel; it affects real training. vLLM 0.20/sm_120 therefore provides no state-training coverage with this mitigation; its new CI result is pending.

The suite retains native comparisons and shares K=128 between generic and Kimi KDA to reuse compilation. The workflow raises only the vLLM 0.30 timeout from 15 to 25 minutes. Setup CODEOWNER review and a green CI run are still required.

KDA bit-exact qualification covers the tested vLLM 0.30 kernel profiles/shapes only. This does not establish full-engine scheduling parity, model-quality recovery or training speed. The private ReplaySSM profile remains outside public CI coverage.

Before your PR is "Ready for review"

  • Backward compatibility: checkpoints missing newly added disabled state-quantizer handles remain loadable, with or without saved linear-attention policy metadata. Enabled W quantization is rejected with a clear error. State QAT requires an explicit serving policy and phase.
  • Dependencies/copying: native vLLM/ReplaySSM kernels are imported. FLA supplies training kernels; the vLLM GPU session installs fla-core==0.5.1 without replacing vLLM's TileLang/TVM pins.
  • Tests: focused CPU policy/QDQ/checkpoint tests, native serving comparisons, and Megatron training/checkpoint restore. Compilation remains in setup fixtures. Core builds without the native gate/packed-forward hooks run a restore/rejection/disabled-baseline check and skip the unsupported full QAT test.
  • Changelog: updated for state QAT.
  • Review: human CODEOWNER approval remains required.

Additional Information

Step 2/6. Known limits include per-layer boundary synchronization, an additional FLA prefix forward for gradients, per-token native cache allocations, and explicit valid lengths for padded-only packed metadata. The source size is intentionally deferred for this review.

Summary by CodeRabbit

  • New Features
    • Added experimental recurrent-state quantization for GDN and KDA, with vLLM serving support for prefill and decode.
    • Added dynamic INT8 state recipes, including block-32 and ReplaySSM/Hadamard options, plus a Kimi-K3 serving profile.
    • Added configurable linear-attention policies and training-phase handling for sequence lengths and state handoff.
  • Documentation
    • Updated recipe guidance with serving requirements, supported profiles, and quantization behavior.
    • Clarified optional serving-kernel requirements and compatibility limitations.

@copy-pr-bot

copy-pr-bot Bot commented Sep 23, 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 23, 2026 •

Copy link
Copy Markdown
Contributor

Review in Change Stack →

Note

Reviews paused

It looks like this branch is under active development. To avoid overwhelming you with review comments due to an influx of new commits, CodeRabbit has automatically paused this review. You can configure this behavior by changing the reviews.auto_review.auto_pause_after_reviewed_commits setting.

Use the following commands to manage reviews:

  • @coderabbitai resume to resume automatic reviews.
  • @coderabbitai review to trigger a single review.

Use the checkboxes below for quick actions:

  • ▶️ Resume reviews
  • 🔍 Trigger review

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration
  • Configuration used: Repository: NVIDIA/Model-Optimizer/.coderabbit.yaml
  • Review profile: CHILL
  • Plan: Enterprise
  • Run ID: d9fff262-effd-42d4-b57f-51349d5837b3






📥 Commits

Reviewing files that changed from the base of the PR and between 2c90acd and fec406f.







📒 Files selected for processing (28)
  • CHANGELOG.rst
  • examples/megatron_bridge/quantize.py
  • modelopt/torch/kernels/quantization/linear_attention/__init__.py
  • modelopt/torch/kernels/quantization/linear_attention/serving/__init__.py
  • modelopt/torch/kernels/quantization/linear_attention/serving/_compat.py
  • modelopt/torch/kernels/quantization/linear_attention/serving/forward.py
  • modelopt/torch/quantization/conversion.py
  • modelopt/torch/quantization/linear_attention/_vllm_autograd.py
  • modelopt/torch/quantization/linear_attention/config.py
  • modelopt/torch/quantization/linear_attention/decode.py
  • modelopt/torch/quantization/linear_attention/gdn.py
  • modelopt/torch/quantization/linear_attention/kda.py
  • modelopt/torch/quantization/linear_attention/training.py
  • modelopt/torch/quantization/linear_attention/utils.py
  • modelopt/torch/quantization/plugins/gdn.py
  • modelopt/torch/quantization/plugins/kda.py
  • modelopt/torch/quantization/plugins/linear_attention.py
  • modelopt/torch/quantization/plugins/megatron.py
  • modelopt_recipes/general/ptq/linear_attention_state_int8_block32_dynamic.yaml
  • modelopt_recipes/general/ptq/linear_attention_state_int8_block32_dynamic_kimi_k3.yaml
  • modelopt_recipes/ptq.md
  • tests/examples/megatron_bridge/test_quantize_export.py
  • tests/gpu_megatron/torch/quantization/plugins/test_megatron_gated_delta_net.py
  • tests/gpu_megatron/torch/quantization/plugins/test_megatron_kda.py
  • tests/gpu_vllm/torch/quantization/test_linear_attention_replay.py
  • tests/gpu_vllm/torch/quantization/test_linear_attention_training.py
  • tests/unit/torch/quantization/plugins/test_gdn.py
  • tests/unit/torch/quantization/test_linear_attention_decode.py






🚧 Files skipped from review as they are similar to previous changes (3)
  • modelopt/torch/kernels/quantization/linear_attention/serving/init.py
  • CHANGELOG.rst
  • modelopt/torch/kernels/quantization/linear_attention/init.py






Included review availability: This review used your included allowance. Your plan provides up to 12 included reviews per hour; 11 remain after this review.








📝 Walkthrough
📝 Walkthrough
📝 Walkthrough
📝 Walkthrough
📝 Walkthrough
📝 Walkthrough

Walkthrough

The change adds serving-aligned GDN and KDA recurrent-state QAT, vLLM forwards, differentiable training and decode paths, policy and checkpoint support, Megatron integration, and dynamic INT8 PTQ recipes. It also updates related tests, documentation, and changelog entries.

Changes

Serving-aligned linear-attention state QAT

Layer / File(s) Summary
Policy configuration and persistence
modelopt/torch/quantization/config.py, modelopt/torch/quantization/linear_attention/config.py, modelopt/torch/quantization/plugins/linear_attention.py, modelopt/torch/quantization/conversion.py, modelopt/torch/quantization/model_quant.py, modelopt/torch/quantization/plugins/__init__.py, modelopt/torch/quantization/linear_attention/__init__.py, modelopt/torch/opt/plugins/mcore_dist_checkpointing.py
QuantizeConfig accepts linear-attention policies. Conversion and restoration apply and validate policies, and checkpoint metadata stores or removes resolved policy state.
Serving kernels and differentiable adapters
modelopt/torch/kernels/quantization/linear_attention/serving/*, modelopt/torch/kernels/quantization/linear_attention/__init__.py, modelopt/torch/quantization/linear_attention/_vllm_autograd.py
vLLM compatibility helpers select kernels and state layouts. Serving forwards implement prefill and decode; the differentiable adapters use native forward results with gradient paths for training. ReplaySSM helpers manage checkpoints and recurrent-kernel calls.
Recurrent state and decode
modelopt/torch/quantization/linear_attention/decode.py, modelopt/torch/quantization/linear_attention/utils.py, tests/unit/torch/quantization/test_linear_attention_decode.py, tests/unit/torch/quantization/test_linear_attention_hadamard.py, tests/gpu_vllm/torch/quantization/test_linear_attention_replay.py
Recurrent decoding adds resumable carries, replay entries, state encoding, and checkpoint refresh. State QDQ supports FP8 and INT8 formats. Unit and GPU tests cover quantization, configuration checks, Hadamard behavior, ReplaySSM output, carry state, and gradients.
Training phases and GDN/KDA integration
modelopt/torch/quantization/linear_attention/{training,gdn,kda}.py, modelopt/torch/quantization/plugins/{gdn,kda,megatron}.py, tests/unit/torch/quantization/plugins/test_gdn.py, tests/gpu_megatron/torch/quantization/plugins/*, tests/gpu_vllm/torch/quantization/test_linear_attention_training.py
Training phases provide per-sequence lengths and packed-boundary metadata. GDN and KDA adapters route supported calls through the shared prefill/decode path. Megatron integration handles serving inputs, selective recomputation, and execution validation. Tests cover policy restoration, supported modes, outputs, state, and gradients.
PTQ recipes and execution constraints
modelopt_recipes/configs/ptq/units/*, modelopt_recipes/general/ptq/linear_attention_state_int8*.yaml, modelopt_recipes/ptq.md, examples/megatron_bridge/quantize.py, tests/examples/megatron_bridge/test_quantize_export.py, CHANGELOG.rst
Recipes add dynamic INT8 state configurations, including block-32 and Kimi-K3 profiles. The Megatron example rejects calibration recipes that include linear-attention policies and skips generation when state quantization is enabled. Documentation and the changelog describe the profiles and requirements.

Priority: ➖ Normal

Estimated code review effort: 5 (Critical) | ~120 minutes

Change: Feature

Sequence Diagram(s)

sequenceDiagram
  participant TrainingPhase as linear_attention_training_phase
  participant Megatron as Megatron linear-attention adapter
  participant QAT as gdn_state_qat or kda_state_qat
  participant Forward as _prefill_decode_forward
  TrainingPhase->>Megatron: set prefill and sequence lengths
  Megatron->>QAT: route supported GDN or KDA call
  QAT->>Forward: pass prepared inputs and state options
  Forward-->>QAT: return output and final state
Loading
















Merge Risk: ⚪ Minimal · up to fec40

This change adds GDN/KDA recurrent-state QAT with serving-aligned forwards and INT8 state recipes. The concerns raised earlier about checkpoint restore and configuration handling have been addressed, and no outstanding defect is established at this head. The change appears ready to merge, subject to the runtime qualification limits the PR already documents.

Pre-merge checks | Passed 5 | Failed 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage Warning Docstring coverage is 42.86% which is insufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 168 functions across 34 files. (4 skipped… Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (5 passed)
Check name Status Explanation
Linked Issues check Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check Passed Check skipped because no linked issues were found for this pull request.
Security Anti-Patterns Passed PASS. The authoritative PR diff adds no listed security anti-pattern in modelopt or examples Python code. It adds no torch.load(..., weights_only=False), numpy.load/np.load(..., allow_pickle=True), ha…
Description Check Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check Passed The title clearly summarizes the main change: Torch GDN/KDA decode QAT with INT8 recurrent-state support. It is concise and directly matches the pull request objectives and changes.

Full details: Docstring Coverage

Explanation

Docstring coverage is 42.86% which is insufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 168 functions across 34 files. (4 skipped: 4 unsupported.)


  • Fix all pre-merge checks with AI
✨ Finishing Touches 💡 3
📝 Generate docstrings 💡
  • Commit to this branch
  • Create a new PR






⚔️ Resolve merge conflicts 💡
  • Resolve merge conflict in branch kaix/linear-attention-decode-first






🧪 Generate unit tests (beta)
  • Commit to this branch
  • Create a new PR







🛠️ Fix failing CI checks 💡
  • Commit to this branch
  • Create a new PR











  • Autopilot · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

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

@kaix-nv kaix-nv changed the title Add decode-first GDN/KDA QAT with INT8 recurrent state [2/4] GDN/KDA decode QAT with INT8 recurrent state Sep 23, 2026
@kaix-nv
kaix-nv added this pull request to stack #2521 September 23, 2026 01:22
@codecov

codecov Bot commented Sep 23, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 75.42628% with 245 lines in your changes missing coverage. Please review.
✅ Project coverage is 78.11%. Comparing base (4c1f813) to head (03f6076).
⚠️ Report is 2 commits behind head on main.

Files with missing lines Patch % Lines
...lopt/torch/quantization/linear_attention/decode.py 63.10% 69 Missing ⚠️
modelopt/torch/quantization/plugins/megatron.py 55.03% 67 Missing ⚠️
...ls/quantization/linear_attention/serving/replay.py 0.00% 52 Missing ⚠️
...s/quantization/linear_attention/serving/forward.py 82.19% 13 Missing ⚠️
modelopt/torch/quantization/plugins/kda.py 38.88% 11 Missing ⚠️
...pt/torch/quantization/linear_attention/training.py 91.86% 10 Missing ⚠️
...s/quantization/linear_attention/serving/_compat.py 90.36% 8 Missing ⚠️
...elopt/torch/quantization/linear_attention/utils.py 91.42% 6 Missing ⚠️
modelopt/torch/quantization/plugins/gdn.py 82.85% 6 Missing ⚠️
...opt/torch/quantization/plugins/linear_attention.py 97.33% 2 Missing ⚠️
... and 1 more
Additional details and impacted files
@@            Coverage Diff             @@
##             main    #2519      +/-   ##
==========================================
+ Coverage   71.76%   78.11%   +6.34%     
==========================================
  Files         644      659      +15     
  Lines       71477    72634    +1157     
==========================================
+ Hits        51298    56740    +5442     
+ Misses      20179    15894    -4285     
Flag Coverage Δ
examples-diffusers 20.20% <18.35%> (-0.04%) ⬇️
examples-gpt-oss 13.49% <15.84%> (+0.02%) ⬆️
examples-hf_ptq 23.25% <17.15%> (-0.13%) ⬇️
examples-llm_distill 13.57% <15.84%> (+0.01%) ⬆️
examples-llm_eval 17.26% <17.15%> (-0.02%) ⬇️
examples-llm_qat 17.58% <18.35%> (-0.01%) ⬇️
examples-llm_sparsity 15.80% <15.84%> (-0.01%) ⬇️
examples-megatron_bridge 26.47% <29.18%> (-0.14%) ⬇️
examples-specdec_bench 13.27% <15.84%> (+0.02%) ⬆️
examples-speculative_decoding 17.55% <17.15%> (-0.10%) ⬇️
examples-torch_onnx 21.56% <17.15%> (-0.08%) ⬇️
examples-torch_trt 15.26% <17.15%> (+0.01%) ⬆️
examples-vllm_serve 13.94% <15.84%> (+0.01%) ⬆️
gpu 59.21% <70.81%> (+25.56%) ⬆️
regression 15.10% <15.84%> (-0.01%) ⬇️
unit 59.49% <32.09%> (-0.39%) ⬇️

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 force-pushed the kaix/linear-attention-decode-first branch from 94080d3 to 757f337 Compare September 23, 2026 19:34
@kaix-nv kaix-nv changed the title [2/4] GDN/KDA decode QAT with INT8 recurrent state [2/5] GDN/KDA decode QAT with INT8 recurrent state 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 removed this pull request from stack #2542 September 24, 2026 06:31
@kaix-nv
kaix-nv added this pull request to stack #2543 September 24, 2026 06:31
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-decode-first branch from 757f337 to 492db57 Compare September 24, 2026 18:17
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-decode-first branch from 492db57 to 5fbb898 Compare September 25, 2026 01:54
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-decode-first branch 3 times, most recently from 5e548c1 to ab35f1e Compare September 25, 2026 20:57
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-decode-first branch 4 times, most recently from 30e0659 to 7b5caf1 Compare September 28, 2026 06:05
@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 [2/7] Torch GDN/KDA decode QAT with INT8 recurrent state [2/6] Torch GDN/KDA decode QAT with INT8 recurrent state Oct 8, 2026
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-decode-first branch from 94c5ead to 2c90acd Compare October 8, 2026 22:40
kaix-nv added a commit that referenced this pull request Oct 10, 2026
Consolidate the unpublished PR #2519 review fixes. Preserve valid packed lengths through selective recompute, expose the public phase API, and keep the Bridge workflow integration separate. Retain the existing FLA kernels for their follow-up PR.

Signed-off-by: Kai Xu <kaix@nvidia.com>
@kaix-nv
kaix-nv requested a review from a team as a code owner October 10, 2026 00:51
kaix-nv added 20 commits October 9, 2026 19:40
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>
Use registered state quantizers for format and last-axis grouping, with INT8 groups of 16, 32, or 64 independent of execution tile width. Preserve existing tile and Hadamard codecs and document the blockwise configuration.

Keep minimal output/gradient, configuration, and checkpoint coverage; omit redundant mocked block routing and quantizer call-count checks.

Signed-off-by: Kai Xu <kaix@nvidia.com>
Make recurrent_decode the single Torch implementation. Process each prepared prefix directly and remove duplicate packing, shape preparation, and the obsolete helper without changing supported numerical policies.

Validated 40 focused CPU tests, the existing Bridge QAT/QAD GPU smoke tests, and 12 before/after output, state, and gradient comparisons.

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>
Unify prefix, decode, and replay execution settings in LinearAttentionConfig with migration for supported nested configs and saved policy objects. Keep formats and grouping in TensorQuantizer.

Retire W-only and reference training paths, move mathematical references under tests, and remove the copied FLA W-QAT implementation and unused recipes.

Validation: 44 focused CPU tests, 6 native and Megatron GPU tests, legacy checkpoint loading, recipe validation, and pre-commit hooks passed.
Signed-off-by: Kai Xu <kaix@nvidia.com>
Normalize recurrent working values to FP32 and require an explicit serving policy for both GDN and KDA state quantization. Remove unpublished config migrations, unused replay-factor plumbing, and the unreferenced Triton INT8 helper while retaining native INT8/Hadamard replay.

Align GDN/KDA API and runtime-state names, document the state handoff, and keep shared TensorQuantizer/FP8 corrections outside this PR. Update recipe guidance and retain minimal real-path tests.

Validation: 28 focused CPU tests, 5 native GPU cases, 2 Megatron QAT/sharded-restore tests, scoped pre-commit hooks, and git diff --check passed.
Signed-off-by: Kai Xu <kaix@nvidia.com>
Signed-off-by: Kai Xu <kaix@nvidia.com>
Adapt native imports, state layouts, recurrent indexing, and KDA gate arithmetic while keeping ModelOpt state QDQ in canonical layout. Use the vllm profile name and retain vllm_0_15 as a legacy alias. Keep lazy gate discovery outside compiled graphs.

Validate policy matches across distributed stages and correct the optional dependency message. Verified 94 CPU tests, native GDN/KDA GPU checks on vLLM 0.15.1, 0.20.0, and 0.30.0, ReplaySSM cases, Megatron backward/checkpoint restore, and pre-commit hooks. Full serving-engine and model-quality qualification remain separate.

Signed-off-by: Kai Xu <kaix@nvidia.com>
Cache state-layout and kernel-signature checks at first use so Sphinx can import serving adapters with mocked optional dependencies. Preserve runtime dispatch and kernel arithmetic.

Validation: focused recursive Sphinx HTML build using repository configuration and warnings as errors; native metadata and CPU state-layout checks on vLLM 0.15.1, 0.20.0, and 0.30.0; scoped pre-commit and git diff --check.
Signed-off-by: Kai Xu <kaix@nvidia.com>
Consolidate the unpublished PR #2519 review fixes. Preserve valid packed lengths through selective recompute, expose the public phase API, and keep the Bridge workflow integration separate. Retain the existing FLA kernels for their follow-up PR.

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-first branch from fec406f to 140735b Compare October 10, 2026 04:40
…lation

Signed-off-by: Kai Xu <kaix@nvidia.com>
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