Skip to content

[3/7] Megatron Bridge linear attention QAT and QAD example - #2657

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

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

Conversation

@kaix-nv

@kaix-nv kaix-nv commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor

Linear-attention series — 7 PRs (1 merged, 6 open)

Order PR Depends on
1/7 #2497 GDN state/W QAT foundation main
2/7 #2519 Torch GDN/KDA decode QAT + INT8 #2497
3/7 #2657 Megatron Bridge linear attention QAT/QAD example #2519
4/7 #2562 Fused Triton GDN/KDA decode QAT #2657
5/7 #2541 vLLM GDN/KDA state-only fake quantization #2562
6/7 #2503 GDN/KDA prefill GEMM quantization #2541
7/7 #2507 Experimental GDN/KDA approximate inverse #2503

The 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. #2657 is rebased onto the latest #2519. Rebase the remaining descendants after their 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 example.

Add a Megatron Bridge example for recurrent-state QAT and QAD. The training entry point selects a native chunked prefix and recurrent suffix, applies loss to the suffix, and keeps the phase context active through backward. Bridge owns optimization, distributed scheduling, and checkpoints. A frozen unquantized teacher enables QAD.

The example contains the training script, dependency files, launcher, two minimal training tests, and concise usage/alignment documentation. Raw experiment records and development links are kept outside the release PR. State formats, execution policies, and kernels are supplied by the parent PR.

Usage

bash examples/llm_qat/linear_attention/with_vllm_defaults.sh \
  torchrun --standalone --nproc-per-node=1 examples/llm_qat/linear_attention/train.py \
  --model /path/to/local-model \
  --train-data /path/to/tokenized/train_text_document \
  --output /path/to/megatron-qat-checkpoint \
  --recipe general/ptq/linear_attention_state_int8_block32_dynamic \
  --train-steps 1 --length 128 --prefill-tokens 64

Add --teacher-model /path/to/unquantized-model for QAD. The ordinary INT8 recipe uses public vLLM kernels. INT8 + Hadamard and replay require a compatible native ReplaySSM fork. A KDA trainer additionally requires a compatible Bridge provider. TP/PP/EP options follow Bridge; context parallelism stays at one, and each local pipeline chunk must contain linear attention.

Testing

After rebasing onto the latest #2519:

  • pytest tests/examples/megatron_bridge/test_linear_attention.py: 2 passed in 99.85 seconds on one RTX A6000, using public vLLM 0.15.1 kernels and matching Megatron Bridge/Core dependencies. Shared setup took 84.22 seconds; QAT and QAD calls each took 4.07 seconds. The tests check student updates and a frozen unquantized QAD teacher.
  • Scoped pre-commit hooks and git diff --check passed. Relative documentation links/anchors and Python snippet syntax were checked.
  • All six example commits replayed without patch changes. The parent-relative diff is 8 files, 821 additions, with no runtime/kernel changes.

Previous workflow checks exercised the real GDN Bridge training entry point with tiny random weights and mock data, including Hadamard token mode and replay QAT/QAD. KDA checks exercised Megatron layers rather than a Bridge KDA trainer. These do not establish pretrained-model quality recovery, full serving-engine equivalence, or performance.

Before your PR is "Ready for review"

  • Is this change backward compatible?: Additive example. The native state-QAT runtime and checkpoint policy are defined by the parent PR.
  • If you copied code or added a dependency, did you follow CONTRIBUTING.md?: The optional public vLLM dependency and its license are listed separately. INT8/Hadamard and replay require the documented native fork.
  • Did you write any new necessary tests?: One QAT and one QAD training-step test with shared compilation setup. This cleanup adds no tests.
  • Did you update Changelog?: N/A; the runtime feature entry belongs to the parent PR.
  • Did you get Claude approval on this PR?: Pending review of the updated example.

Additional Information

Step 3/7. Merge #2519 first, then this example, followed by #2562. This example trains state quantization; prefill GEMM quantization and approximate inverse remain follow-up work. Rebased onto #2519 at 94c5ead61c; the review diff contains only this example (8 files, 821 added lines).

@copy-pr-bot

copy-pr-bot Bot commented Oct 5, 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 Oct 5, 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
  • 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 [2/2] Megatron Bridge linear attention QAT and QAD example Megatron Bridge linear attention QAT and QAD example (companion to #2519) Oct 5, 2026
@kaix-nv
kaix-nv added this pull request to stack #2563 October 5, 2026 06:11
@kaix-nv
kaix-nv removed this pull request from stack #2563 October 5, 2026 06:15
@kaix-nv
kaix-nv added this pull request to stack #2658 October 5, 2026 06:15
@kaix-nv kaix-nv changed the title Megatron Bridge linear attention QAT and QAD example (companion to #2519) [3/7] Megatron Bridge linear attention QAT and QAD example Oct 5, 2026
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-example branch from 2829885 to 3837d28 Compare October 6, 2026 17:27
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-example branch from 3837d28 to b7401f7 Compare October 6, 2026 19:55
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-example branch from b7401f7 to 0357915 Compare October 6, 2026 21:26
Signed-off-by: Kai Xu <kaix@nvidia.com>
Signed-off-by: Kai Xu <kaix@nvidia.com>
Update the linear-attention example for the flat execution config, native state-only training, and checkpoint migration. Remove the retired CPU reference script while retaining its historical results and source link.

Record the focused CPU, native GPU, Megatron checkpoint, and config migration validation. Documentation pre-commit hooks passed.

Signed-off-by: Kai Xu <kaix@nvidia.com>
Document the QDQ scheduling and arithmetic mismatch, the native prefix/handoff/suffix solution, gradient flow, and the limits of current validation. Align configuration guidance with the current state QAT API.

Signed-off-by: Kai Xu <kaix@nvidia.com>
Remove raw experiment JSON from the example and shorten the usage and state-alignment documents. Retain the mismatch derivation, supported policies, concise numerical results, and validation limits. Preserve development evidence in an untracked local archive.

Training code and minimal QAT/QAD tests are unchanged. Validation: scoped pre-commit, documentation links and anchors, Python snippet syntax, archive integrity, and git diff --check.

Signed-off-by: Kai Xu <kaix@nvidia.com>
@kaix-nv
kaix-nv force-pushed the kaix/linear-attention-qat-example branch from 0d898ac to 7717ea4 Compare October 8, 2026 05:35
@codecov

codecov Bot commented Oct 8, 2026

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 78.62%. Comparing base (94c5ead) to head (7717ea4).

Additional details and impacted files
@@                         Coverage Diff                         @@
##           kaix/linear-attention-decode-first    #2657   +/-   ##
===================================================================
  Coverage                               78.62%   78.62%           
===================================================================
  Files                                     650      650           
  Lines                                   71262    71262           
===================================================================
  Hits                                    56029    56029           
  Misses                                  15233    15233           
Flag Coverage Δ
unit 60.03% <ø> (ø)

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.

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