Repository navigation
[1/7] GDN state/W QAT foundation - #2497
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. |
|
Navigate logical layers of code changes, visualize relationships, and explore their blast radius. Note Reviews pausedIt 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 Use the following commands to manage reviews:
Use the checkboxes below for quick actions:
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configurationConfiguration used: Repository: NVIDIA/Model-Optimizer/.coderabbit.yaml Review profile: CHILL Plan: Enterprise Run ID: 📒 Files selected for processing (2)
Included review availability: This review used your included allowance. Your plan provides up to 12 included reviews per hour; 11 remain after this review. 📝 WalkthroughWalkthroughThis change adds experimental dynamic FP8 fake quantization for GatedDeltaNet recurrent states and WY activations. It adds chunked Triton kernels, quantizer and Megatron integration, PTQ configurations, and reference and GPU tests. The documented fused path uses ChangesGatedDeltaNet Quantization
Priority: ⬇️ Low Estimated code review effort: 4 (Complex) | ~60 minutes Change: Feature Merge Risk: ⚪ Minimal · up to The change reorganizes test setup so compilation is warmed before functional tests. No actionable merge-blocking risk was found. 🚥 Pre-merge checks | ✅ 5 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (5 passed)
✨ Finishing Touches 💡 1📝 Generate docstrings 💡
🧪 Generate unit tests (beta)
Comment |
|
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## main #2497 +/- ##
==========================================
+ Coverage 69.49% 78.34% +8.85%
==========================================
Files 614 618 +4
Lines 68631 69455 +824
==========================================
+ Hits 47694 54414 +6720
+ Misses 20937 15041 -5896
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:
|
bff7ddf to
6686c9b
Compare
| chunks.append((n, qc, kc, gc, decay, u)) | ||
| all_w.append(w.transpose(0, 1)) | ||
| w = torch.cat(all_w).reshape(q.shape) | ||
| if w_quantizer is not None: |
|
After we QAT a model with this PR, how do we evaluate it? Do you have a corresponding vLLM patch? |
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>
Signed-off-by: Kai Xu <kaix@nvidia.com>
| from modelopt.torch.kernels.quantization.common.fp8_quant import fp8_scalar_qdq | ||
|
|
||
| # ``STATE_QDQ`` modes of the forward state kernel. | ||
| STATE_QDQ_OFF = 0 |
There was a problem hiding this comment.
Is this how we are controlling the behavior of the QAT, through these global variables? Is there a way to make this more programatic?
There was a problem hiding this comment.
These constants identify kernel mode. The QAT config is controlled per module through ModelOpt’s quant_cfg.
|
|
||
| # ``STATE_QDQ`` modes of the forward state kernel. | ||
| STATE_QDQ_OFF = 0 | ||
| STATE_QDQ_FP8_DYNAMIC = 1 # FP8 E4M3, one dynamic scale per program tile ([K, BV] of one head) |
There was a problem hiding this comment.
Can we also add support for integer state?
There was a problem hiding this comment.
Added dynamic INT8 state fake quantization by moving the basic implementation and tests from PR #2519 into this PR.
Centralize the pinned FLA test dependencies in the dev-fla optional extra used by both GPU nox sessions. Replace the exhaustive compilation and numerical matrices with three shared-shape BF16 checks and one small Megatron QAT/checkpoint test; compile only the selected paths in setup fixtures. Consolidates the dependency, extra-naming, and minimal-test follow-ups without adding Triton/TileLang CI cache plumbing. On RTX A6000, cold and warm focused runs each passed 3 tests with 1 SM89+ hardware skip; the warm run took 38.20s total and 2.41s in functional calls. Dependency wiring and pre-commit checks passed. Signed-off-by: Kai Xu <kaix@nvidia.com>
7236936 to
7b43e6e
Compare
| - example: gpu | ||
| timeout: 60 | ||
| # Includes dependency builds and cold compilation of the FLA forward/backward tests. | ||
| timeout: 75 |
There was a problem hiding this comment.
@kaix-nv gpu tests now only take 40mins after your simplifications. Can we revert this change and leave it as 60mins now?
There was a problem hiding this comment.
Addressed in d4895f8353: restored the general GPU job timeout to 60 minutes and removed the comment added with the increase. Workflow YAML validation and pre-commit checks passed.
Signed-off-by: Kai Xu <kaix@nvidia.com>
Preserve both GDN and upstream indexer entries in the changelog and PTQ unit table. Keep the reviewed GDN implementation, minimal GPU tests, and 60-minute GPU job timeout. Validation: 36 focused GDN/reference unit tests passed; pre-commit and conflict checks passed. Signed-off-by: Kai Xu <kaix@nvidia.com>
_QuantGatedDeltaNet (#2497) gets the class-level get/set_extra_state overrides that torch needs to route its quantizer state through _extra_state; GatedDeltaNet defines none, so without them the state was dropped, as test_registered_megatron_quant_modules_checkpoint_quantizer_state flags. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
cjluo-nv
left a comment
There was a problem hiding this comment.
Bot review (bedrock-claude-opus-5-5) — DM the bot to share feedback.
Nudge: the code-level concerns from the last review are fixed, but the PR is still 2000 core-logic lines against a 500-line budget. It also vendors MIT-licensed FLA code, which needs a human licensing sign-off.
Needs action:
-
✂️ Split this PR into stacked
[x/N]PRs. Each PR must build and pass CI on its own and carry its own tests. Suggested order:[1/4]:plugins/gated_delta_net.py,linear_attention/utils.py, theconversion.py/model_quant.pyhooks, the recipe units, and the unit and reference tests.[2/4]:fla_chunk_delta_h.py.[3/4]:fla_chunk_gated_delta_rule.pywith its GPU test.[4/4]: the Megatron adapter inplugins/megatron.py.
The kernel directory (~1723 lines) is still over budget on its own. Consider landing an unmodified upstream copy first, then the
[ModelOpt]edits. Renumber the 7-PR stack to match. -
Confirm third-party approval with a human. That covers the vendored
fla-corekernels, theLICENSEcopyright entry, theApache-2.0 AND MITSPDX headers and the newdev-flapins (fla-core,tilelang,apache-tvm-ffi). The PR body says approval is still pending. -
Confirm where INT8 state support lives. A reply says it was moved into this PR, but the diff only implements FP8, and
test_validate_state_quantizer_rejects_unsupportedrejects INT8. -
Fix the PR body: the title says
[1/7]and the URL is #2497, but the body says #2497 is already merged.
No action needed:
- ✔️ Resolved since the last review:
- Kernel detection now checks identity against FLA's function and unwraps
functools.partial. - Leftover kwargs now raise
TypeError. - The CI timeout, workflow and Mamba-timeout edits were reverted.
- The dependencies moved to the
dev-flaextra. - Kernel compilation moved into fixtures.
- Kernel detection now checks identity against FLA's function and unwraps
Linear-attention series — 7 PRs (1 merged, 6 open)
mainThe 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. 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. #2506 and #2509 are superseded and closed.
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_cfginterface and start disabled. State QDQ uses 64-token chunks and recomputesamaxat 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_qdqapplies that supplied scale rather than choosing tensor-wide grouping. W grouping is applied byTensorQuantizer. Quantizer settings use normal ModelOpt checkpoint state. There is noQuantizeConfig.linear_attentionfield 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 inlinear_attention/utils.py.The fused path requires
fla-core==0.5.1and 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
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-flaoptional 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:
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.
CONTRIBUTING.md: ❌ Internal third-party approval tracking still needs confirmation. Upstream attribution, MIT/Apache headers,LICENSEnotice, and license-hook exclusions are included. FLA/TileLang and TVM-FFI license files were reviewed.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.
Summary by CodeRabbit
fla-core==0.5.1and a chunk size of 64.