diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index f60fad5cc15..a5a3b97cd8c 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -112,8 +112,7 @@ repos: modelopt/torch/_deploy/utils/onnx_utils.py| modelopt/torch/export/transformer_engine.py| modelopt/torch/fastgen/plugins/qwen_image_pdd.py| - modelopt/torch/kernels/quantization/linear_attention/fla_chunk_delta_h.py| - modelopt/torch/kernels/quantization/linear_attention/fla_chunk_gated_delta_rule.py| + modelopt/torch/kernels/quantization/linear_attention/serving/chunk_delta_h.py| modelopt/torch/puzzletron/anymodel/models/gpt_oss/gpt_oss_pruned_to_mxfp4.py| modelopt/torch/quantization/export_onnx.py| modelopt/torch/quantization/plugins/attention.py| diff --git a/CHANGELOG.rst b/CHANGELOG.rst index a6bc486f252..44a8540e99a 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -13,9 +13,12 @@ Changelog *Quantization* +- Add experimental FP8/INT8 recurrent-state fake quantization before native vLLM GDN/KDA prefill and decode; use the state-only recipe with the documented vLLM 0.15.x eager runtime. + - Add Q8_0 weight-only quantization with 32-value GGML blocks, packed unified HF and Megatron export, and a built-in ``q8_0`` PTQ recipe. - Add Hugging Face PTQ calibration and export support for Nemotron-H MTP modules, preserving calibrated expert input scales. - Backfill checkpoint aliases for three more published NVFP4 releases: ``moonshotai/Kimi-K2.7-Code``, ``google/gemma-4-26B-A4B-it`` and ``google/diffusiongemma-26B-A4B-it``. Each imports an existing general recipe wholesale rather than copying its body. +- Add experimental serving-aligned GDN/KDA state QAT with explicit prefill/decode boundaries; migrate chunk-only state recipes to a serving policy and supply prefix lengths through the training phase context. Plain state QDQ uses the installed vLLM's native kernels via ``precision="vllm"`` (the legacy ``vllm_0_15`` spelling remains accepted); INT8 Hadamard and ReplaySSM require the compatible quantized-ReplaySSM serving fork. - Backfill checkpoint aliases for nine more published NVFP4 releases, so each is reachable from its source model's hub path: ``zai-org/GLM-5.1`` and ``GLM-5.2``, ``MiniMaxAI/MiniMax-M2.5`` and ``MiniMax-M3``, ``deepseek-ai/DeepSeek-V3.1`` and ``DeepSeek-V3.2``, ``Qwen/Qwen3-235B-A22B-Instruct-2507`` and ``-Thinking-2507``, and ``Qwen/Qwen3.6-27B``. Each imports an existing general or architecture recipe wholesale rather than copying its body. - Add composed Hugging Face AutoQuantize recipes that run fixed PTQ or weight AutoQuantize before a separate KV-cache AutoQuantize stage, with independent resumable checkpoints for the weight and @@ -32,7 +35,6 @@ Changelog - Add support for quantizing and calibrating enabled operators outside the transformer layers, such as ``lm_head``, when using layerwise calibration. - Add an end-to-end BEVFormer ONNX PTQ example with temporal calibration data generation, INT8 and FP8 quantization, TensorRT engine building, and nuScenes accuracy evaluation. See `examples/onnx_ptq/bevformer/README.md `_ for details. - Add a reusable local-Hessian NVFP4 PTQ recipe and the quantization recipe used for ``nvidia/Qwen3.8-27B-NVFP4``. -- Add experimental dynamic FP8 fake quantization of GatedDeltaNet chunk-boundary states and WY activations for training through the standard ``quant_cfg`` interface. The fused GDN path requires ``fla-core==0.5.1`` and chunk size 64; state emulation requires SM89 or newer. - Add fake quantization of the sparse-attention indexer key cache and query for DeepSeek-V4 (vLLM and Megatron-Core) and GLM-5.3-Flash (vLLM) through the new ``indexer_k_quantizer`` and ``indexer_q_quantizer``. Enable them by importing the ``configs/ptq/units/indexer_k_nvfp4`` and ``configs/ptq/units/indexer_q_nvfp4`` units (NVFP4 with the global scale fixed to 1) into a recipe. - Add the ``configs/ptq/units/kv_nvfp4_mla`` recipe unit for fake quantization of the MLA KV cache (DeepSeek-V3, GLM-5.3-Flash, ...) in vLLM fake-quant serving: NVFP4 for the latent and FP8 for the RoPE key. See `examples/vllm_serve/README.md `_ for an example recipe. - vLLM fake-quant serving now runs on pre-quantized checkpoints such as FP8 when the recipe leaves those layers unquantized (for example a KV-cache-only recipe), and on MLA models with an FP8 KV cache; both previously failed during quantization. @@ -65,6 +67,8 @@ Changelog **Backward Breaking Changes** - Bump minimum transformers version to ``5.5`` instead of ``4.57``; transformers 4.x is no longer supported. Upgrade with ``pip install -U "nvidia-modelopt[hf]"``. +- Experimental GDN W quantization and its ``gdn_w_fp8_dynamic`` recipe unit are removed; keep ``gdn_w_quantizer`` disabled when using state QAT or restoring existing checkpoints. + - ``modelopt.torch.distill.plugins.megatron.TopKLogitsKLLoss`` (``logit_kl_topk`` in ``DistillationConfig``) is renamed to ``TopLogitsKLLoss``, keeping the old name as a deprecated alias. It now normalizes both distributions over the full vocabulary instead of re-normalizing over the Top-K entries, and always appends a "ghost" token holding the probability mass outside the Top-K to both student and teacher (matching Megatron-LM's offline cached-logits KD loss). Loss values change for existing ``logit_kl_topk`` runs. - ``LogitsAndIntermediatesLossBalancer`` (Megatron distillation plugin) no longer rescales the distillation loss to the magnitude of the LM loss when the LM loss is included. The total is now the fixed convex combination ``(1 - alpha) * lm_loss + alpha * kd_loss`` with ``DistillationConfig.kd_loss_alpha`` (in [0, 1]), matching Megatron-LM's offline cached-logits KD. The default ``kd_loss_alpha=1.0`` skips the LM loss, as before. ``DistillationConfig.skip_lm_loss`` and ``kd_loss_scale`` are removed and raise ``ValueError`` if passed; set ``kd_loss_alpha`` instead. In ``examples/megatron_bridge/distill.py``, ``--kd_loss_alpha`` replaces ``--no_skip_lm_loss`` and ``--kd_loss_scale``. - The ``examples/vllm_serve`` fakequant launcher no longer supports vLLM 0.9.0. Upgrade to a version listed as tested in the example README. diff --git a/LICENSE b/LICENSE index abc6a0cb761..0e99a82684b 100644 --- a/LICENSE +++ b/LICENSE @@ -216,6 +216,7 @@ Portions of this repository were adapted from code originally authored by the following copyright holders, licensed under the Apache License, Version 2.0 (see full license text above): + Copyright contributors to the vLLM project Copyright 2021 The HuggingFace Inc. team Copyright 2022 The HuggingFace Team Copyright 2022, Lefebvre Dalloz Services diff --git a/examples/llm_qat/README.md b/examples/llm_qat/README.md index cd88ae30498..a32c25368d7 100644 --- a/examples/llm_qat/README.md +++ b/examples/llm_qat/README.md @@ -13,6 +13,7 @@ For background on QAT and QAD and help choosing between Hugging Face, Megatron B | Arguments | Full CLI/YAML argument reference | \[[Link](ARGUMENTS.md)\] | | Support Matrix | Supported models, quantization formats, and backends | \[[Link](#support-matrix)\] | | QLoRA | Model training with reduced GPU memory | \[[Link](#qlora-real-quantization)\] | +| Linear Attention | GDN/KDA recurrent-state QAT and ReplaySSM example | \[[Link](linear_attention/README.md)\] | | Advanced Topics | Trainer APIs, FSDP2 config, YAML options | \[[Link](#advanced-topics)\] | | Results | Accuracy benchmarks | \[[Link](#results)\] | | Resources | Extra links and references | \[[Link](#resources)\] | diff --git a/examples/llm_qat/linear_attention/README.md b/examples/llm_qat/linear_attention/README.md new file mode 100644 index 00000000000..776b8287861 --- /dev/null +++ b/examples/llm_qat/linear_attention/README.md @@ -0,0 +1,216 @@ +# Quantization-Aware Training and Distillation for Linear Attention + +This example trains Megatron-Core GDN or KDA attention parameters with recurrent +state fake quantization. Megatron Bridge runs the optimizer, distributed training, +and checkpointing. QAT uses next-token cross-entropy; QAD uses a frozen, +unquantized teacher and Bridge's logits distillation loss. + +Training combines native chunked prefill with a recurrent suffix that follows +the selected serving arithmetic and state quantization schedule. Loss is applied +to the suffix, with gradients flowing through the state handoff into the prefix. +See [State quantization alignment](STATE_QUANTIZATION.md) for the mismatch, +solution, and numerical validation. The example quantizes recurrent state only. + +## Requirements + +Use the [Megatron Bridge environment](../../megatron_bridge/README.md#pre-requisites), +CUDA GPUs, and mutually compatible Bridge/Core revisions. ModelOpt adapts +Megatron `GatedDeltaNet` and `KimiDeltaAttention`; FLA supplies their kernel +dependency. KDA requires a Core revision exporting `KimiDeltaAttention` and a +Bridge provider that can convert the chosen checkpoint. Layer support in +ModelOpt alone does not provide checkpoint conversion support in Bridge. + +Prepare a local model/tokenizer snapshot and a Megatron `.bin`/`.idx` dataset +using that tokenizer. `--train-data` takes the shared filename prefix without an +extension; see [data preparation](../../megatron_bridge/README.md#data-preparation). +For a model already built in Megatron, use the Python API below. + +For ordinary state QDQ, install the example dependencies from the repository root: + +```bash +pip install -r examples/llm_qat/linear_attention/requirements.txt +pip install -r examples/llm_qat/linear_attention/requirements-vllm.txt +``` + +The example pins `fla-core==0.5.1` and public `vllm==0.15.1`. vLLM supplies forward +kernels; no running server is needed for training. INT8 + Hadamard and ReplaySSM +require a compatible quantized-ReplaySSM fork instead of public vLLM, including +its KDA vector-gate kernel when training KDA. + +## Run QAT or QAD + +```bash +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. Student and teacher must +have matching tokenizer vocabularies and output vocabulary dimensions. QAD uses +`kd_loss_alpha=1.0`, disabling the language-model loss. Both QAT and QAD mask loss +to positions after `--prefill-tokens`. + +Only the student's linear-attention parameters are trainable; the remaining +parameters are frozen. Training uses BF16 mixed precision. Each dense sequence +has the same fixed prefix length in this CLI; packed data or variable boundaries +require a custom batch integration. + +`with_vllm_defaults.sh` sets `FLA_USE_FAST_OPS=0`, `USE_DEFAULT_FLA_NORM=0`, +`FLA_GDN_FIX_BT=0`, `FLA_USE_CUDA_GRAPH=0`, and `FLA_TRIL_PRECISION=ieee` before +Python imports vLLM. Use the same runtime and settings for training and serving. +For evaluation, wrap the server or worker launch; wrapping a client does not +configure an already-running server. On multiple nodes, apply the wrapper to +workers on every node. The library does not enforce these settings at import. + +## Select a state quantization recipe + +`--recipe` accepts a built-in name or a custom YAML path. The complete recipes +under `general/ptq/` configure both GDN and KDA state quantizers and their +execution policy. They use dynamic quantization without a calibration pass. + +| Recipe | State quantization | Runtime | +| --- | --- | --- | +| `linear_attention_state_int8_block32_dynamic` (default) | INT8 QDQ at handoff and every suffix token; one scale per key row and 32 value channels | Public vLLM | +| `linear_attention_state_int8_dynamic` | INT8 with Hadamard rotation over 32 value channels; checkpoint every suffix token | Compatible quantized-ReplaySSM fork | +| Same Hadamard recipe with `replay_window=8` | Checkpoint at handoff and every eight suffix tokens; BF16 key/update ring between refreshes | Compatible quantized-ReplaySSM fork | + +Hadamard requires the value dimension to be divisible by 32. The state-only vLLM +fake-quant adapter targets ordinary TensorQuantizer QDQ; the Hadamard/replay +recipes target the separate native ReplaySSM implementation. + +ModelOpt separates quantization settings from execution and batch metadata: + +| Setting | Purpose | Where it belongs | +| --- | --- | --- | +| `TensorQuantizer` | Enable state QDQ; select format and scale grouping | Recipe `quant_cfg` | +| `LinearAttentionConfig` | Select native arithmetic and checkpoint frequency | Recipe `linear_attention` rules | +| Prefix lengths | Specify where each sequence switches to recurrence | `linear_attention_training_phase` at runtime | + +For example, a complete recipe can select this execution policy: + +```yaml +linear_attention: + - module_name: "*" + cfg: + backend: serving + precision: replayssm + replay_window: 8 +``` + +`precision="vllm"` selects the installed public vLLM arithmetic; `vllm_0_15` is +an accepted legacy spelling. `precision="replayssm"` selects native INT8/Hadamard +checkpoints. `replay_window=1` refreshes every token; values from 2 to 64 require +ReplaySSM. The prefill chunk size of 64 is independent of this suffix window. + +`TensorQuantizer.block_sizes` owns blockwise scale grouping. `state_block_v` +controls the execution tile width and legacy per-tile grouping; it does not +replace `block_sizes`. An execution policy alone does not enable a quantizer. +Importing only a recipe unit configures the quantizers without its full policy. + +To change the Hadamard checkpoint window in Python: + +```python +from modelopt.recipe import load_recipe + +cfg = load_recipe("general/ptq/linear_attention_state_int8_dynamic").quantize.model_dump() +cfg["linear_attention"][0]["cfg"]["replay_window"] = 8 +``` + +For the CLI, put the modified recipe in a YAML file and pass its path to +`--recipe`. State QAT requires `backend="serving"` and explicit prefix lengths. +Keep the legacy `gdn_w_quantizer` disabled; this workflow does not quantize W. + +## Supply prefill/decode boundaries in a training loop + +A 128-token sequence with prefix length 96 runs tokens 0–95 through native +chunked prefill, encodes the handoff, then runs tokens 96–127 recurrently. The +prefix contains a full 64-token chunk and a partial 32-token chunk; the phase +switch happens at token 96, independently of the chunk boundaries. + +The CLI supplies one fixed length for every sequence. The API accepts different +lengths per sequence, such as `[64, 96]`. These are batch metadata, so they are not +saved as part of the quantization policy. Keep the context active through +backward when activation checkpointing recomputes the forward. + +```python +import torch + +import modelopt.torch.quantization as mtq +from modelopt.recipe import load_recipe +from modelopt.torch.quantization.linear_attention import linear_attention_training_phase + +# model is an initialized Megatron model; ids and shifted labels have shape [2, 128]. +cfg = load_recipe("general/ptq/linear_attention_state_int8_block32_dynamic").quantize +model = mtq.quantize(model, cfg) +optimizer = torch.optim.AdamW(model.parameters(), lr=1e-5) +model.train() +optimizer.zero_grad(set_to_none=True) +positions = torch.arange(128, device=ids.device).expand_as(ids) +with linear_attention_training_phase(model, [64, 96]): + with torch.autocast("cuda", dtype=torch.bfloat16): + losses = model( + input_ids=ids, position_ids=positions, attention_mask=None, labels=labels + ) + mask = torch.arange(128, device=ids.device)[None, :] >= torch.tensor( + [64, 96], device=ids.device + )[:, None] + loss = (losses * mask).sum() / mask.sum() + loss.backward() +optimizer.step() +``` + +`[0]` means recurrence only; `[T]` means an all-prefix sequence and performs no +suffix handoff QDQ. QAT/QAD needs a nonempty suffix to train against state +rounding. The context selects execution phases; the caller still supplies +labels and loss masking. It restores the previous lengths on exit. +Concurrent forwards with different phase contexts on the same model are unsupported. + +## Checkpointing and multiple GPUs + +Bridge saves model, optimizer/scheduler, and ModelOpt state under +`/checkpoints`. Reusing `--output` resumes from the latest checkpoint; +increase `--train-steps` to the desired total step count. Retain the model, +teacher, topology, dataset, and prefix boundary. The restored checkpoint supplies +its saved quantization policy; prefix lengths must still be supplied at runtime. +Export to a serving model is a separate step. + +Use `--tp_size`, `--pp_size`, and `--ep_size` as in the +[Megatron Bridge example](../../megatron_bridge/distill.py). Student and teacher +use the same topology. Sequence parallelism is enabled when TP exceeds one; +additional ranks use data parallelism. `--global-batch-size` controls accumulation +with microbatch size one and must be divisible by the data-parallel size. + +Choose a topology supported by the model's Bridge/Core provider. This example +requires linear-attention layers in each local pipeline chunk. Context parallelism +is fixed at one. The available topology options do not imply multi-GPU qualification. + +## Validation and limitations + +The minimal example tests run one GDN QAT step and one QAD step, with shared +compilation setup. They check student weight updates and a frozen, unquantized +QAD teacher. Run them in the matching Bridge/vLLM environment: + +```bash +bash examples/llm_qat/linear_attention/with_vllm_defaults.sh \ + python -m pytest tests/examples/megatron_bridge/test_linear_attention.py +``` + +GDN Bridge workflow checks and KDA Megatron layer checks are distinct: the tested +Bridge provider supports GDN, while a KDA trainer needs a compatible provider. +Native-cache and gradient checks are summarized in +[State quantization alignment](STATE_QUANTIZATION.md#validation-results). + +The recurrent suffix uses Python orchestration and a Torch adjoint, so long +suffixes can be slow. Kernel/cache agreement does not establish pretrained-model +quality recovery, full serving-engine equivalence, or training speed. Prefill +GEMM quantization and approximate inverse are outside this example. + +## State-only serving with vLLM + +The [vLLM example](../../vllm_serve/README.md#linear-attention-state-quantization) +uses the same state QDQ helper for plain FP8/INT8 policies. Match the runtime, +quantizer grouping, and scheduler prefill boundaries to the training setup. diff --git a/examples/llm_qat/linear_attention/STATE_QUANTIZATION.md b/examples/llm_qat/linear_attention/STATE_QUANTIZATION.md new file mode 100644 index 00000000000..4840c27fe59 --- /dev/null +++ b/examples/llm_qat/linear_attention/STATE_QUANTIZATION.md @@ -0,0 +1,161 @@ +# State quantization alignment between training and serving + +GDN/KDA state QAT must reproduce the state consumed by serving. This requires +matching the quantization boundaries, scale groups, native arithmetic, and output +read timing. The implementation combines a native chunked prefix with a recurrent +suffix and supplies a differentiable Torch adjoint for training. + +## Why chunk-only quantization does not match decode + +For one head, the recurrent state has shape `[K, V]`: key channels by value +channels. Each token updates this state and reads an output. Quantize/dequantize +(QDQ) rounds the state; the training API retains floating-point values. + +Let `F_t` be token `t`'s update and `Q` be QDQ. Quantizing once at a chunk boundary +and quantizing every decode state generally give different results: + +$$ +Q(F_2(F_1(S))) \ne Q(F_2(Q(F_1(S)))). +$$ + +The second decode update consumes the first update's rounded state. Chunk-only +QDQ omits that intermediate perturbation. For a toy update `S = S + 0.6`, starting +at zero and rounding to the nearest integer: + +| Event | QDQ after two tokens | QDQ between tokens | +| --- | --- | --- | +| First working state | 0.6 | 0.6 | +| State consumed by token 2 | 0.6 | 1.0 | +| Second working state | 1.2 | 1.6 | +| Rounded state for the next token | 1 | 2 | + +This is an illustration of rounding placement, not the actual INT8 scaling rule. +A straight-through estimator (STE) changes the backward approximation; it cannot +repair this forward mismatch. + +Matching the QDQ schedule alone is insufficient. BF16 casts, reduction order, +normalization, and gate evaluation can change values near quantization thresholds. +ReplaySSM also reconstructs state from a checkpoint and a weighted BF16 update +ring. An algebraically equivalent FP32 recurrence can follow a different rounded +trajectory. Training therefore imports the selected serving forward kernels. + +## Match each training phase to serving + +All tokens come from the dataset through teacher forcing. A per-sequence boundary +selects native prefill followed by native recurrence: + +```text +prompt tokens completion tokens +[native chunked prefill] -> encoded handoff -> [native token/replay updates] + training prefix training suffix +``` + +| Phase | Training behavior | Serving behavior being reproduced | +| --- | --- | --- | +| Fresh prefix | Native chunked prefill without internal state QDQ | One native prompt-prefill call | +| Continuation prefix | Encode its incoming nonzero state | Quantize the state supplied to a continuation-prefill call | +| Handoff | Encode final prefix state when the suffix is nonempty | Quantize incoming state before the first decode token | +| Suffix | Native updates with token writes or replay checkpoint refreshes | The selected decode/cache implementation | + +The prefix stays chunked because it represents serving prefill. Adding QDQ after +every 64-token training chunk would introduce boundaries absent from a single +native prefill call. A prompt split across multiple serving prefill calls requires +matching each actual incoming-state boundary; one prefix length alone does not +encode the scheduler's call sequence. + +`linear_attention_training_phase(model, prefill_lengths)` supplies the boundary. +Keep it active through forward and backward so activation recomputation sees the +same split. The example masks QAT/QAD loss to the suffix; gradients remain +connected through the handoff into the prefix. An all-prefix sequence has no +suffix handoff QDQ and supplies no suffix training signal. + +## Why state QDQ before and after an update can align + +For ordinary token-state QDQ, training stores rounded state after each update; +a state-only serving adapter can round its floating cache before the next call. +These placements align when they apply the same deterministic QDQ once to the +same working state. The current output is read before checkpoint rounding. + +Let `S_P` be the prefix's final state, `C_t` the training carry, `R_t` the serving +cache, and `U_t` the working state. Training computes: + +$$ +C_0 = Q(S_P),\qquad +U_t = F_t(C_{t-1}),\qquad +o_t = \operatorname{read}_t(U_t),\qquad +C_t = Q(U_t). +$$ + +Serving computes: + +$$ +R_0 = S_P,\qquad +U'_t = F_t(Q(R_{t-1})),\qquad +o'_t = \operatorname{read}_t(U'_t),\qquad +R_t = U'_t. +$$ + +Initially `C_0 = Q(R_0)`. If the token updates consume equal states, identical +native arithmetic gives `U_t = U'_t`, equal outputs, and `C_t = Q(R_t)`. The +invariant therefore holds for the next token. Raw stored caches need not match; +the states consumed by the recurrence must match. + +This argument requires matching inputs, initial states, QDQ format/grouping, +runtime settings, and update/readout arithmetic. It does not justify applying +QDQ twice or comparing different kernel specializations. + +## INT8, Hadamard, and ReplaySSM + +Ordinary state QDQ uses TensorQuantizer. The block32 INT8 recipe uses one scale +per key row and 32 value channels and encodes at handoff and every suffix token. + +The native ReplaySSM profile instead rotates groups of 32 value channels into a +Hadamard basis, then scales, rounds to INT8, and reconstructs the checkpoint. +Hadamard rotation and quantization are separate operations. Checkpoints use a +scale per key row and 32-value group, with FP16 scale metadata. + +With `replay_window=1`, every token refreshes the checkpoint. Larger windows keep +that checkpoint and append each native key/update vector once in BF16, then +refresh at the window boundary. GDN uses scalar decay; KDA's channel-dependent +decay acts on the key axis while Hadamard rotates the value axis. The update ring, +rounding, and checkpoint schedule all need to match serving. + +The adapter imports the native encoder and recurrent kernels from a compatible +quantized-ReplaySSM fork. These interfaces, including KDA vector-gate support, +are separate from the public vLLM profile and ordinary state-only fake quantization. +Temporary native cache buffers protect tensors needed by backward. The persistent +training carry remains floating fake-QDQ data. + +The Torch adjoint uses saved forward values and identity STE through casts and +QDQ. This is a training approximation to rounding, with gradients through the +prefix, handoff, token writes, and checkpoint refreshes. It is not an exact +derivative of the discontinuous quantizer. + +## Validation results + +The following checks used RTX A6000 GPUs. The replay and training matrix used +Torch 2.9.1, FLA 0.5.1, public vLLM 0.15.1 prefill, and compatible native ReplaySSM +sources with local KDA extensions. Results describe the tested implementation +snapshots and environments. + +| Check | Result | +| --- | --- | +| 22 GDN/KDA native replay cases, including windows 1/4/16, nonzero states, partial prefixes, longer sequences, and QDQ-off controls | Exact outputs, checkpoint values/scales, and BF16 ring entries; finite gradients into nonempty prefixes | +| Four split/resume cases with Q/K normalization | Exact outputs, final states, and gradients; empty calls perform no write | +| Four GDN Bridge runs: QAT/QAD with Hadamard token mode and replay window 4 | One-step student updates; QAD teachers remain frozen and unquantized | +| Four KDA Megatron layer runs with training or distillation losses | Student updates and gradients into the prefix; separate from Bridge trainer validation | +| Public-vLLM GDN/KDA checks on 0.15.1, 0.20.0, and 0.30.0 | Exact native output/final-state comparisons, QDQ placement/effect, and gradients in the focused kernel tests | + +The Bridge checks used the example's training entry point, tiny random weights, +and mock data, including activation recomputation and the distributed optimizer. +The tested Bridge provider supported GDN; KDA requires its own compatible +provider. Matching versions and arithmetic settings remains necessary even where +more than one runtime version passes kernel tests. + +These checks resolve the reproduced kernel/cache mismatch for the tested cases. +They do not establish full-model serving equivalence: model conversion, +projections, convolution, normalization, engine-managed multi-call prefill, +cache-slot reuse, and scheduler batching still need matched integration tests. +Pretrained-model quality recovery, multi-GPU scaling, and training throughput +also require separate validation. Long recurrent suffixes can be slow with the +current Python orchestration and Torch backward implementation. diff --git a/examples/llm_qat/linear_attention/requirements-vllm.txt b/examples/llm_qat/linear_attention/requirements-vllm.txt new file mode 100644 index 00000000000..e725f5d767d --- /dev/null +++ b/examples/llm_qat/linear_attention/requirements-vllm.txt @@ -0,0 +1,3 @@ +# Optional native arithmetic profile; kernels are imported without starting a server. +# Apache-2.0: https://github.com/vllm-project/vllm/blob/v0.15.1/LICENSE +vllm==0.15.1 diff --git a/examples/llm_qat/linear_attention/requirements.txt b/examples/llm_qat/linear_attention/requirements.txt new file mode 100644 index 00000000000..915edf246cd --- /dev/null +++ b/examples/llm_qat/linear_attention/requirements.txt @@ -0,0 +1 @@ +fla-core==0.5.1 diff --git a/examples/llm_qat/linear_attention/train.py b/examples/llm_qat/linear_attention/train.py new file mode 100644 index 00000000000..eb725156604 --- /dev/null +++ b/examples/llm_qat/linear_attention/train.py @@ -0,0 +1,263 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Run Megatron Bridge QAT or QAD with GDN/KDA recurrent-state fake quantization.""" + +import argparse +from contextlib import ExitStack +from pathlib import Path + +import torch +from megatron.bridge import AutoBridge +from megatron.bridge.models.distillation_provider import convert_to_distillation_provider +from megatron.bridge.training.config import ( + CheckpointConfig, + ConfigContainer, + DistributedDataParallelConfig, + GPTDatasetConfig, + LoggerConfig, + OptimizerConfig, + RNGConfig, + SchedulerConfig, + TokenizerConfig, + TrainingConfig, + ValidationConfig, +) +from megatron.bridge.training.distill import distill +from megatron.bridge.training.gpt_step import forward_step_modelopt +from megatron.bridge.training.post_training.checkpointing import ( + has_modelopt_state, + load_modelopt_state, +) +from megatron.bridge.training.post_training.distillation import ModelOptDistillConfig +from megatron.bridge.training.pretrain import pretrain +from megatron.bridge.training.state import GlobalState +from megatron.bridge.utils.vocab_utils import calculate_padded_vocab_size +from megatron.core.models.common.embeddings.language_model_embedding import LanguageModelEmbedding +from megatron.core.ssm import gated_delta_net +from megatron.core.utils import unwrap_model +from transformers import AutoTokenizer + +import modelopt.torch.quantization as mtq +from modelopt.recipe import ModelOptPTQRecipe, load_recipe +from modelopt.torch.quantization.linear_attention import linear_attention_training_phase +from modelopt.torch.quantization.utils import is_quantized + + +def run_training(config, quant_config, prefill_tokens, teacher_provider=None): + """Delegate optimization and checkpointing to Bridge, with a fixed phase per sequence.""" + if not 0 <= prefill_tokens < config.dataset.seq_length: + raise ValueError("prefill-tokens must leave at least one suffix label") + resume = config.checkpoint.load and has_modelopt_state(config.checkpoint.load) + with ExitStack() as phases: + + def prepare_student(models): + # Restore the saved policy before quantization and before the QAD conversion hook. + if resume: + load_modelopt_state(models, config.checkpoint.load) + layer_types = (gated_delta_net.GatedDeltaNet,) + if hasattr(gated_delta_net, "KimiDeltaAttention"): + layer_types += (gated_delta_net.KimiDeltaAttention,) + for student in unwrap_model(models): + layers = [m for m in student.modules() if isinstance(m, layer_types)] + if not layers: + raise ValueError( + "Each local model chunk must contain Megatron GatedDeltaNet or " + "KimiDeltaAttention; choose a pipeline layout with linear attention on every stage" + ) + student.requires_grad_(False) + for layer in layers: + layer.requires_grad_(True) + if config.model.recompute_granularity == "full": + + def require_input_grad(module, args, output): + return output.requires_grad_(True) if module.training else output + + # Reentrant checkpointing needs a differentiable input with frozen embeddings. + for module in student.modules(): + if isinstance(module, LanguageModelEmbedding): + handle = module.register_forward_hook(require_input_grad) + phases.callback(handle.remove) + if not is_quantized(student): + mtq.quantize(student, quant_config) + phases.enter_context( + linear_attention_training_phase( + student, [prefill_tokens] * config.train.micro_batch_size + ) + ) + return models + + def masked_batch(data_iterator): + batch = dict(next(data_iterator)) + batch["loss_mask"] = batch["loss_mask"].clone() + batch["loss_mask"][..., :prefill_tokens] = 0 + yield batch + + def forward_step(state: GlobalState, data_iterator, model, return_schedule_plan=False): + # Consume lazily: middle pipeline stages may not need any batch tensors. + return forward_step_modelopt( + state, masked_batch(data_iterator), model, return_schedule_plan + ) + + config.model.register_pre_wrap_hook(prepare_student) + if teacher_provider is not None: + config.model = convert_to_distillation_provider( + config.model, teacher_provider, ModelOptDistillConfig(kd_loss_alpha=1.0) + ) + distill(config, forward_step) + else: + pretrain(config, forward_step) + + +def model_provider(path, options, *, load_weights=True): + """Build a Megatron provider with the same topology options as the Bridge example.""" + bridge = AutoBridge.from_hf_pretrained(str(path), trust_remote_code=options.trust_remote_code) + provider = bridge.to_megatron_provider(load_weights=load_weights) + provider.tensor_model_parallel_size = options.tp_size + provider.pipeline_model_parallel_size = options.pp_size + provider.pipeline_dtype = torch.bfloat16 + provider.context_parallel_size = 1 + provider.expert_model_parallel_size = options.ep_size + provider.expert_tensor_parallel_size = 1 + provider.sequence_parallel = options.tp_size > 1 + provider.gradient_accumulation_fusion = False + provider.calculate_per_token_loss = True + provider.seq_length = options.length + return provider + + +def main(): + """Train a local model with QAT, or add a frozen teacher for QAD.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model", type=Path, required=True) + parser.add_argument("--teacher-model", type=Path, help="Unquantized teacher; enables QAD") + parser.add_argument( + "--train-data", + type=Path, + required=True, + help="Megatron tokenized dataset prefix (.bin/.idx)", + ) + parser.add_argument( + "--recipe", + default="general/ptq/linear_attention_state_int8_block32_dynamic", + help="Path to a quantization recipe YAML (built-in or custom)", + ) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--trust-remote-code", action="store_true") + parser.add_argument("--train-steps", type=int, default=1) + parser.add_argument("--length", type=int, default=128) + parser.add_argument("--prefill-tokens", type=int, default=64) + parser.add_argument("--seed", type=int, default=2026) + parser.add_argument("--learning-rate", type=float, default=1e-5) + parser.add_argument("--global-batch-size", type=int, default=1) + parser.add_argument("--tp_size", type=int, default=1, help="Tensor parallel size") + parser.add_argument("--pp_size", type=int, default=1, help="Pipeline parallel size") + parser.add_argument("--ep_size", type=int, default=1, help="Expert parallel size") + options = parser.parse_args() + if options.train_steps < 1 or options.length < 2: + parser.error("train-steps must be positive and length must be at least two") + if not 0 <= options.prefill_tokens < options.length: + parser.error("prefill-tokens must leave at least one suffix label") + if options.global_batch_size < 1: + parser.error("global-batch-size must be positive") + if min(options.tp_size, options.pp_size, options.ep_size) < 1: + parser.error("Parallel sizes must be positive") + if any(not Path(f"{options.train_data}.{suffix}").is_file() for suffix in ("bin", "idx")): + parser.error("train-data must name an existing Megatron .bin/.idx dataset prefix") + + recipe = load_recipe(options.recipe) + if not isinstance(recipe, ModelOptPTQRecipe): + parser.error("--recipe must select a PTQ quantization recipe") + + checkpoint_dir = str(options.output / "checkpoints") + resume = has_modelopt_state(checkpoint_dir) + student = model_provider(options.model, options, load_weights=not resume) + teacher = None + if options.teacher_model is not None: + tokenizer_kwargs = {"trust_remote_code": options.trust_remote_code} + student_vocab = AutoTokenizer.from_pretrained(options.model, **tokenizer_kwargs).get_vocab() + teacher_vocab = AutoTokenizer.from_pretrained( + options.teacher_model, **tokenizer_kwargs + ).get_vocab() + if student_vocab != teacher_vocab: + parser.error("QAD requires student and teacher to use the same tokenizer vocabulary") + teacher = model_provider(options.teacher_model, options) + vocab_sizes = [ + calculate_padded_vocab_size( + p.vocab_size, p.make_vocab_size_divisible_by, p.tensor_model_parallel_size + ) + for p in (student, teacher) + ] + if vocab_sizes[0] != vocab_sizes[1]: + parser.error("QAD requires matching student and teacher output vocabulary dimensions") + + config = ConfigContainer( + model=student, + train=TrainingConfig( + train_iters=options.train_steps, + global_batch_size=options.global_batch_size, + micro_batch_size=1, + ), + validation=ValidationConfig(eval_iters=0, eval_interval=options.train_steps), + optimizer=OptimizerConfig( + optimizer="adam", + lr=options.learning_rate, + min_lr=0, + weight_decay=0, + clip_grad=1.0, + use_distributed_optimizer=True, + ), + scheduler=SchedulerConfig( + lr_decay_style="constant", + lr_warmup_iters=0, + start_weight_decay=0, + end_weight_decay=0, + use_checkpoint_opt_param_scheduler=True, + ), + ddp=DistributedDataParallelConfig( + average_in_collective=False, use_distributed_optimizer=True + ), + dataset=GPTDatasetConfig( + seq_length=options.length, + blend=([str(options.train_data)], None), + split="100,0,0", + random_seed=options.seed, + reset_position_ids=False, + reset_attention_mask=False, + eod_mask_loss=False, + dataloader_type="single", + num_workers=0, + ), + tokenizer=TokenizerConfig( + tokenizer_type="HuggingFaceTokenizer", + tokenizer_model=str(options.model), + hf_tokenizer_kwargs={"trust_remote_code": options.trust_remote_code}, + ), + checkpoint=CheckpointConfig( + save=checkpoint_dir, + load=checkpoint_dir, + save_interval=options.train_steps, + async_save=False, + ckpt_format="torch_dist", + ), + logger=LoggerConfig(log_interval=1), + rng=RNGConfig(seed=options.seed), + mixed_precision="bf16_mixed", + ) + run_training(config, recipe.quantize, options.prefill_tokens, teacher) + + +if __name__ == "__main__": + main() diff --git a/examples/llm_qat/linear_attention/with_vllm_defaults.sh b/examples/llm_qat/linear_attention/with_vllm_defaults.sh new file mode 100644 index 00000000000..6e91436dc64 --- /dev/null +++ b/examples/llm_qat/linear_attention/with_vllm_defaults.sh @@ -0,0 +1,36 @@ +#!/usr/bin/env bash + +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +set -euo pipefail + +if [[ $# -eq 0 ]]; then + echo "Usage: bash $0 [args...]" >&2 + exit 2 +fi + +# Establish the same arithmetic before training or serving imports vLLM. +for setting in \ + FLA_USE_FAST_OPS=0 \ + USE_DEFAULT_FLA_NORM=0 \ + FLA_GDN_FIX_BT=0 \ + FLA_USE_CUDA_GRAPH=0 \ + FLA_TRIL_PRECISION=ieee; do + export "$setting" + printf '%s\n' "$setting" >&2 +done + +exec "$@" diff --git a/examples/vllm_serve/README.md b/examples/vllm_serve/README.md index 8eaea64700d..eebef04de4d 100644 --- a/examples/vllm_serve/README.md +++ b/examples/vllm_serve/README.md @@ -9,6 +9,10 @@ compact NVFP4 attention worker documented below requires vLLM 0.15.0 or newer. The fakequant launcher does not support vLLM 0.9.0. Use one of the tested releases above. +For GDN/KDA recurrent-state QDQ before native prefill and decode, see the +[linear-attention example](#linear-attention-state-quantization), which +documents the initial vLLM 0.15.x runtime requirements and state-only INT8 recipe. + ## Prepare environment Run the commands below from the ModelOpt repository root (`/workspace/Model-Optimizer` @@ -405,6 +409,95 @@ Supported configurations are regular decoder self-attention with FlashInfer or F Unsupported features are sliding window, ALiBi, softcap, sinks, FP8 KV cache, cross/encoder/MLA attention, KV sharing or transfer, prefix caching, speculative decoding, DBO/ubatching, and `FULL` mixed/prefill CUDA graphs. +## Linear-attention state quantization + +The existing `FakeQuantWorker` adds a ModelOpt `TensorQuantizer` immediately before +native vLLM prefill and decode for `Qwen3NextGatedDeltaNet` and +`KimiDeltaAttention`. Native attention kernels, projections, convolution, gates, +outputs, and cache management remain in use. + +### Runtime and launch + +The initial adapter targets vLLM 0.15.x V1 with key-first +`[slot, head, Dk, Dv]` FP32 recurrent state. Runtime checks reject other layouts +and versions. Use a vLLM 0.15.x environment for this adapter; the newer +versions supported by the general fakequant example use different model wrappers. + +```bash +PYTHONPATH=.:examples/vllm_serve \ +RECIPE_PATH=examples/vllm_serve/linear_attention_state_int8.yaml \ +bash examples/llm_qat/linear_attention/with_vllm_defaults.sh \ +python examples/vllm_serve/vllm_serve_fakequant.py /path/to/model \ + --tensor-parallel-size 2 --enforce-eager --no-async-scheduling \ + --no-enable-prefix-caching --mamba-cache-dtype float32 +``` + +The example enables signed symmetric dynamic INT8 state QDQ. Import the `configs/ptq/units/gdn_state_fp8_dynamic` quantizer unit for GDN FP8 +E4M3, or set `num_bits: [4, 3]` on the corresponding state quantizers. Each invocation +receives `[active_sequence, local_head, Dk, Dv]`; `axis: [0, 1]` retains one +scale per sequence and local head **within each value-column tile**. The shared +training/serving QDQ helper splits the last axis using `state_block_v` (default +64), then reduces over `[Dk, tile_width]`. State tensors remain FP32. + +For per-key 32-value INT8 blocks, use the existing recipe +`modelopt_recipes/general/ptq/linear_attention_state_int8_block32_dynamic.yaml`. +Its `block_sizes: {-1: 32}` is handled directly by TensorQuantizer and overrides +legacy tile grouping. The same recipe can be used for state QAT. + +Dynamic scales need no calibration dataset; `algorithm: null` skips dataset +loading. Weight/activation calibration can use the worker's existing recipe +and calibration loop. Static state calibration is unsupported in this adapter. + +Use `MODELOPT_STATE_PATH=/path/to/modelopt_state.pt` instead of a recipe to +restore quantizer configuration. Reload maps recurrent-state quantizer names +through the model's HF-to-vLLM mapper. Checkpoint weights must already match +the vLLM model architecture. + +### Quantization boundaries + +- Prefill: vLLM gathers initial states and zeros fresh requests. The wrapper + applies `TensorQuantizer` to this tensor and calls the original prefill kernel. +- Decode: the wrapper gathers active native cache slots, quantizes them, writes + them back, and calls the original decode kernel. Other slots are untouched. +- Native kernels compute outputs and update recurrent state normally. No + additional rounding is applied to the kernel's final-state write. +- The next call quantizes that state before reading it. A fresh zero state is + unchanged; prompt-to-decode rounding happens before the first decode call. +- Scheduler-level chunked prefill creates one QDQ boundary per invocation. + Internal kernel chunks do not create extra state-quantization boundaries. + Results can therefore depend on scheduler prompt-chunk sizes. +- Requests, cache-slot reuse, and preemption remain managed by vLLM. The plugin + allocates temporary gathered states, with no additional persistent state cache. +- TP ranks quantize their local heads independently, without scale all-reduce. + +The adapter accepts plain `precision: vllm` policies, including the legacy +`vllm_0_15` spelling. The serving engine supplies the phase boundaries, so no +training phase context is required. W quantization and the native ReplaySSM/ +Hadamard profile require separate serving integration and are rejected. + +Training and serving share state-QDQ formats and grouping. Numerical alignment +also requires the same vLLM runtime, arithmetic settings, and prefill scheduling: +a training prefix must reproduce each serving prompt-continuation boundary. +The QAT handoff rounds the prefix state before its first decode read; this +adapter rounds that same state before the first native decode call. QAT stores the +rounded next-state checkpoint after each update; serving stores the working +state and rounds it at the next read. These feed the same rounded state to +the next recurrence. Raw terminal cache buffers therefore need not match the +QAT checkpoint until the same QDQ is applied. + +The supported scope is eager synchronous execution with TP and PP=DP=CP=1. +Speculative decoding, prefix caching, state transfer, and CUDA graphs require +separate integration. This is floating-point numerical emulation; model-quality, +capacity, and performance claims require separate measurements. + +### Later prefill GEMM support + +Prefill operand QDQ will be a separate change after an optimized fused kernel +is available. It will reuse the eight numerical sites from the training prefill +implementation and preserve this native-cache wrapper. The PyTorch materialized +backend is not exposed in vLLM by this state-only adapter. State QDQ cadence and +prefill operand QDQ remain independent numerical policies. + ## Known Problems 1. **MCore reload does not use `MODELOPT_STATE_PATH`**; use `QUANT_FILE_PATH` and make sure `QUANT_CFG` matches the quantization recipe used for the original MCore model (otherwise quantizer keys/config won’t align). diff --git a/examples/vllm_serve/fakequant_worker.py b/examples/vllm_serve/fakequant_worker.py index 6bc82fbc22f..e3bbcadc7ba 100644 --- a/examples/vllm_serve/fakequant_worker.py +++ b/examples/vllm_serve/fakequant_worker.py @@ -38,6 +38,7 @@ disable_compilation, post_restore_vllm_parallel_linears, ) +from modelopt.torch.quantization.plugins.vllm_linear_attention import bind_vllm_linear_attention from modelopt.torch.utils import safe_load from modelopt.torch.utils.dataset_utils import get_dataset_dataloader @@ -114,17 +115,22 @@ def _fakequant_run_prolog_worker(self, mlflow_tracker: FakeQuantMlflowTracker) - print("Will load quant, so only do a single sample calibration") quant_config["calib_size"] = 1 - calib_dataloader = get_dataset_dataloader( - dataset_name=quant_config["dataset"], - tokenizer=tokenizer, - batch_size=quant_config["calib_batch_size"], - num_samples=quant_config["calib_size"], - device=self.device, - ) + quant_cfg = get_quant_config(quant_config, model) + calibrate_loop = None + if quant_cfg.get("algorithm", "max") is not None: + calib_dataloader = get_dataset_dataloader( + dataset_name=quant_config["dataset"], + tokenizer=tokenizer, + batch_size=quant_config["calib_batch_size"], + num_samples=quant_config["calib_size"], + device=self.device, + ) + worker_loop = calibrate_fun(calib_dataloader, self) - calibrate_loop = calibrate_fun(calib_dataloader, self) + def calibrate_loop(converted): + bind_vllm_linear_attention(converted, self.model_runner) + worker_loop(converted) - quant_cfg = get_quant_config(quant_config, model) # Before calibration, which is the run this artifact is most wanted for if it dies. mlflow_tracker.log_quant_config(quant_cfg) @@ -134,6 +140,7 @@ def _fakequant_run_prolog_worker(self, mlflow_tracker: FakeQuantMlflowTracker) - quantizer_file_path = quant_config["quant_file_path"] if quantizer_file_path: + bind_vllm_linear_attention(model, self.model_runner) self.model_runner._dummy_run(1) current_state_dict = load_state_dict_from_path(self, quantizer_file_path, model) model.load_state_dict(current_state_dict) @@ -148,6 +155,7 @@ def _fakequant_run_prolog_worker(self, mlflow_tracker: FakeQuantMlflowTracker) - if isinstance(module, TensorQuantizer): module.to(self.device) + bind_vllm_linear_attention(model, self.model_runner) if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0: mtq.print_quant_summary(model) mlflow_tracker.log_quant_summary(model) diff --git a/examples/vllm_serve/linear_attention_state_int8.yaml b/examples/vllm_serve/linear_attention_state_int8.yaml new file mode 100644 index 00000000000..e0784ac64ed --- /dev/null +++ b/examples/vllm_serve/linear_attention_state_int8.yaml @@ -0,0 +1,15 @@ +imports: + base_disable_all: configs/ptq/units/base_disable_all + state_int8: configs/ptq/units/linear_attention_state_int8_dynamic + +metadata: + recipe_type: ptq + description: GDN/KDA INT8 state-tile QDQ before native vLLM prefill and decode. +quantize: + quant_cfg: + - $import: base_disable_all + - $import: state_int8 + linear_attention: + - module_name: '*' + cfg: {backend: serving, precision: vllm} + algorithm: diff --git a/examples/vllm_serve/vllm_reload_utils.py b/examples/vllm_serve/vllm_reload_utils.py index 1bfc7f95fd4..7973d4e71bf 100644 --- a/examples/vllm_serve/vllm_reload_utils.py +++ b/examples/vllm_serve/vllm_reload_utils.py @@ -241,6 +241,19 @@ def _merge_values_require_identical(merged_key: str, key_value_pairs: list[tuple return first_value +def _map_linear_attention_names(values, map_fun): + """Map module paths using weight-path probes without dropping numerical policies.""" + if map_fun is None: + return values + probes = {(name + ".weight" if name else "weight"): value for name, value in values.items()} + mapped = map_fun(probes) + if len(mapped) != len(values) or any( + name != "weight" and not name.endswith(".weight") for name in mapped + ): + raise ValueError("HF-to-vLLM mapping must preserve every linear-attention module") + return {name.removesuffix("weight").removesuffix("."): value for name, value in mapped.items()} + + def convert_dict_to_vllm( state_dict: dict[str, Any], max_or_concat: bool = True, @@ -254,6 +267,15 @@ def convert_dict_to_vllm( max_or_concat: Whether to merge grouped values by taking max/concatenate or require identical map_fun: Function to map the state dict to vLLM format """ + state_dict = dict(state_dict) + linear_quantizers = {} + for key in list(state_dict): + match = re.match(r"^(.*)\.((?:gdn|kda)_(?:state|w)_quantizer(?:\..*)?)$", key) + if match: + module, suffix = match.groups() + mapped = _map_linear_attention_names({module: state_dict.pop(key)}, map_fun) + name, value = next(iter(mapped.items())) + linear_quantizers[f"{name}.{suffix}"] = value # If map_fun is provided, pre-transform quantizer key module-path prefixes so that # HF→vLLM model renames (e.g. backbone.layers → model.layers) are applied before # key grouping (q/k/v → qkv, experts.N.up_proj → experts.w13, etc.). @@ -287,7 +309,7 @@ def convert_dict_to_vllm( _, value = key_value_pairs[0] vllm_state_dict[merged_key] = value if map_fun is None: - return vllm_state_dict + return {**vllm_state_dict, **linear_quantizers} # Quantizer module-path keys (e.g. "layers.0.mlp.gate_proj.input_quantizer") must NOT # go through map_fun (hf_to_vllm_mapper.apply_dict), which maps weight tensor paths and # drops any key it doesn't recognise — including all quantizer keys. Split them out, @@ -295,7 +317,7 @@ def convert_dict_to_vllm( quantizer_keys = {k: v for k, v in vllm_state_dict.items() if "_quantizer" in k} non_quantizer_keys = {k: v for k, v in vllm_state_dict.items() if "_quantizer" not in k} mapped = map_fun(non_quantizer_keys) if non_quantizer_keys else {} - return {**mapped, **quantizer_keys} + return {**mapped, **quantizer_keys, **linear_quantizers} def convert_modelopt_state_to_vllm( @@ -320,6 +342,16 @@ def convert_modelopt_state_to_vllm( modelopt_state_dict = modelopt_state.pop("modelopt_state_dict", []) for idx, current_mode in enumerate(modelopt_state_dict): current_mode_metadata = current_mode[1].pop("metadata", {}) + if "linear_attention" in current_mode_metadata: + current_mode_metadata["linear_attention"] = _map_linear_attention_names( + current_mode_metadata["linear_attention"], map_fun + ) + # Use exact saved module names instead of training-framework selector patterns. + if "quant_cfg" in current_mode[1]["config"]: + current_mode[1]["config"]["linear_attention"] = [ + {"module_name": name, "cfg": policy} + for name, policy in current_mode_metadata["linear_attention"].items() + ] current_mode_quant_state = current_mode_metadata.pop("quantizer_state", {}) if current_mode_quant_state: current_mode_metadata["quantizer_state"] = convert_dict_to_vllm( @@ -367,6 +399,15 @@ def filter_modelopt_state_quantizer_state_for_model( metadata = mode_entry[1].get("metadata", {}) if "quantizer_state" in metadata: saved = metadata["quantizer_state"] + missing_linear = [ + name + for name, state in saved.items() + if re.search(r"(?:gdn|kda)_(?:state|w)_quantizer$", name) + and not state.get("_disabled", False) + and name not in model_keys + ] + if missing_linear: + raise ValueError(f"Saved linear-attention quantizers are missing: {missing_linear}") # Keep keys that exist in the model. Remove disabled quantizers UNLESS they # have registered buffers (e.g. _pre_quant_scale from AWQ/smoothquant on a diff --git a/modelopt/torch/kernels/quantization/linear_attention/__init__.py b/modelopt/torch/kernels/quantization/linear_attention/__init__.py index 3f003f403f8..617bb6dcbd8 100644 --- a/modelopt/torch/kernels/quantization/linear_attention/__init__.py +++ b/modelopt/torch/kernels/quantization/linear_attention/__init__.py @@ -13,12 +13,8 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Linear-attention kernels for quantization. +"""Linear-attention quantization adapters for optional serving kernels. -``fla_chunk_delta_h.py`` and ``fla_chunk_gated_delta_rule.py`` are adapted copies of the chunked -GatedDeltaNet kernels of `flash-linear-attention `_ -(``fla.ops.common.chunk_delta_h`` and ``fla.ops.gated_delta_rule.chunk``) that can fake-quantize the -recurrent state carried between chunks to FP8 (``state_qdq``). They still import the surrounding -fla operators, so ``fla-core==0.5.1`` and Triton must be installed to use -them. This package initializer does not import the kernels, so importing it needs neither. +Serving adapters import forward kernels only when selected; this package +initializer has no kernel dependencies. """ diff --git a/modelopt/torch/kernels/quantization/linear_attention/fla_chunk_delta_h.py b/modelopt/torch/kernels/quantization/linear_attention/fla_chunk_delta_h.py deleted file mode 100644 index 875cb39a7d0..00000000000 --- a/modelopt/torch/kernels/quantization/linear_attention/fla_chunk_delta_h.py +++ /dev/null @@ -1,978 +0,0 @@ -# Adapted from: https://github.com/fla-org/flash-linear-attention/blob/516143e31fce/fla/ops/common/chunk_delta_h.py -# Adapted with modifications (marked [ModelOpt]): optional in-kernel FP8 E4M3 fake quantization -# of the carried chunk state (STATE_QDQ), BV as an explicit launch argument, no fla backend -# dispatch or config-cache autotune. -# -# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li -# -# This source code is licensed under the MIT license found in the -# LICENSE file in the root directory of this source tree. -# For a list of all contributors, visit: -# https://github.com/fla-org/flash-linear-attention/graphs/contributors - - -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 AND MIT -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import torch -import triton -import triton.language as tl -from fla.ops.utils import prepare_chunk_indices, prepare_chunk_offsets -from fla.ops.utils.cache import fla_cache_autotune -from fla.ops.utils.op import exp2 -from fla.utils import ( - IS_INTEL, - IS_NVIDIA_BLACKWELL, - IS_NVIDIA_HOPPER, - autotune_cache_kwargs, - check_shared_mem, -) - -from modelopt.torch.kernels.quantization.common.fp8_quant import fp8_scalar_qdq - -# ``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) -STATE_QDQ_MAX_BLOCK_V = 128 - - -@triton.jit -def _state_qdq_scale(b_h1, b_h2, b_h3, b_h4, K: tl.constexpr): - """[ModelOpt] Dynamic FP8 E4M3 scale per full [K, BV] tile of one sequence and head.""" - b_amax = tl.max(tl.abs(b_h1)) - if K > 64: - b_amax = tl.maximum(b_amax, tl.max(tl.abs(b_h2))) - if K > 128: - b_amax = tl.maximum(b_amax, tl.max(tl.abs(b_h3))) - if K > 192: - b_amax = tl.maximum(b_amax, tl.max(tl.abs(b_h4))) - return tl.where(b_amax > 0, b_amax / 448.0, 1.0) - - -NUM_WARPS = [2, 4] if IS_NVIDIA_HOPPER else [2, 4, 8, 16] - -# TODO: Triton mainline fixes a Blackwell tl.dot recurrence race. -# Keep this kernel on num_warps=2 for Blackwell until Triton 3.8 is released -# and we re-validate the wider config space. -# Intel needs more warps than NVIDIA here: 8 warps is ~1.5x faster than the best -# config reachable under the [2, 4] cap. -if IS_NVIDIA_BLACKWELL: - GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2] -elif IS_INTEL: - GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2, 4, 8, 16] -else: - GATED_DELTA_RULE_FWD_H_NUM_WARPS = [2, 4] - - -@triton.heuristics( - { - "USE_G": lambda args: args["g"] is not None, - "USE_GK": lambda args: args["gk"] is not None, - "USE_INITIAL_STATE": lambda args: args["h0"] is not None, - "STORE_FINAL_STATE": lambda args: args["ht"] is not None, - "SAVE_NEW_VALUE": lambda args: args["v_new"] is not None, - "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, - } -) -# ``BV`` is an explicit argument rather than an autotuned config: with ``STATE_QDQ`` it sets the -# quantization granularity, so it must not vary with the autotuner's choice. -@triton.autotune( - configs=[ - triton.Config({}, num_warps=num_warps, num_stages=num_stages) - for num_warps in GATED_DELTA_RULE_FWD_H_NUM_WARPS - for num_stages in ([2, 3, 4] if check_shared_mem("ampere") else [2, 1]) - ], - key=["H", "HV", "K", "V", "BT", "BV", "STATE_V_FIRST", "STATE_QDQ"], -) -@triton.jit(do_not_specialize=["T"]) -def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( - k, - v, - w, - v_new, - g, - gk, - h, - h0, - ht, - cu_seqlens, - chunk_offsets, - T, - H: tl.constexpr, - HV: tl.constexpr, - K: tl.constexpr, - V: tl.constexpr, - BT: tl.constexpr, - BV: tl.constexpr, - USE_G: tl.constexpr, - USE_GK: tl.constexpr, - USE_INITIAL_STATE: tl.constexpr, - STORE_FINAL_STATE: tl.constexpr, - SAVE_NEW_VALUE: tl.constexpr, - STATE_V_FIRST: tl.constexpr, - IS_VARLEN: tl.constexpr, - STATE_QDQ: tl.constexpr, -): - pid = tl.program_id(0) - NV = tl.cdiv(V, BV) - i_v, i_nh = pid % NV, (pid // NV).to(tl.int64) - i_n, i_h = i_nh // HV, i_nh % HV - if IS_VARLEN: - bos, eos = ( - tl.load(cu_seqlens + i_n).to(tl.int64), - tl.load(cu_seqlens + i_n + 1).to(tl.int64), - ) - T = eos - bos - NT = tl.cdiv(T, BT) - boh = tl.load(chunk_offsets + i_n).to(tl.int64) - else: - bos, eos = i_n * T, i_n * T + T - NT = tl.cdiv(T, BT) - boh = i_n * NT - - if STATE_V_FIRST: - b_h1 = tl.zeros([BV, 64], dtype=tl.float32) - if K > 64: - b_h2 = tl.zeros([BV, 64], dtype=tl.float32) - if K > 128: - b_h3 = tl.zeros([BV, 64], dtype=tl.float32) - if K > 192: - b_h4 = tl.zeros([BV, 64], dtype=tl.float32) - else: - b_h1 = tl.zeros([64, BV], dtype=tl.float32) - if K > 64: - b_h2 = tl.zeros([64, BV], dtype=tl.float32) - if K > 128: - b_h3 = tl.zeros([64, BV], dtype=tl.float32) - if K > 192: - b_h4 = tl.zeros([64, BV], dtype=tl.float32) - - # calculate offset - h += (boh * HV + i_h).to(tl.int64) * K * V - v += (bos * HV + i_h).to(tl.int64) * V - k += (bos * H + i_h // (HV // H)).to(tl.int64) * K - w += (bos * HV + i_h).to(tl.int64) * K - if SAVE_NEW_VALUE: - v_new += (bos * HV + i_h).to(tl.int64) * V - - if USE_INITIAL_STATE: - h0 = h0 + i_nh * K * V - if STORE_FINAL_STATE: - ht = ht + i_nh * K * V - - # load initial state - o_v = i_v * BV + tl.arange(0, BV) - m_v = o_v < V - o_k1 = tl.arange(0, 64) - m_k1 = o_k1 < K - o_k2 = 64 + o_k1 - m_k2 = o_k2 < K - o_k3 = 128 + o_k1 - m_k3 = o_k3 < K - o_k4 = 192 + o_k1 - m_k4 = o_k4 < K - if USE_INITIAL_STATE: - if STATE_V_FIRST: - p_h0_1 = h0 + o_v[:, None] * K + o_k1[None, :] - m_h0_1 = m_v[:, None] & m_k1[None, :] - else: - p_h0_1 = h0 + o_k1[:, None] * V + o_v[None, :] - m_h0_1 = m_k1[:, None] & m_v[None, :] - b_h1 += tl.load(p_h0_1, mask=m_h0_1, other=0.0).to(tl.float32) - if K > 64: - if STATE_V_FIRST: - p_h0_2 = h0 + o_v[:, None] * K + o_k2[None, :] - m_h0_2 = m_v[:, None] & m_k2[None, :] - else: - p_h0_2 = h0 + o_k2[:, None] * V + o_v[None, :] - m_h0_2 = m_k2[:, None] & m_v[None, :] - b_h2 += tl.load(p_h0_2, mask=m_h0_2, other=0.0).to(tl.float32) - if K > 128: - if STATE_V_FIRST: - p_h0_3 = h0 + o_v[:, None] * K + o_k3[None, :] - m_h0_3 = m_v[:, None] & m_k3[None, :] - else: - p_h0_3 = h0 + o_k3[:, None] * V + o_v[None, :] - m_h0_3 = m_k3[:, None] & m_v[None, :] - b_h3 += tl.load(p_h0_3, mask=m_h0_3, other=0.0).to(tl.float32) - if K > 192: - if STATE_V_FIRST: - p_h0_4 = h0 + o_v[:, None] * K + o_k4[None, :] - m_h0_4 = m_v[:, None] & m_k4[None, :] - else: - p_h0_4 = h0 + o_k4[:, None] * V + o_v[None, :] - m_h0_4 = m_k4[:, None] & m_v[None, :] - b_h4 += tl.load(p_h0_4, mask=m_h0_4, other=0.0).to(tl.float32) - # [ModelOpt] A state read from an FP8 cache is quantized before the first chunk uses it. - if STATE_QDQ == 1: - if K > 192: - b_scale = _state_qdq_scale(b_h1, b_h2, b_h3, b_h4, K=K) - elif K > 128: - b_scale = _state_qdq_scale(b_h1, b_h2, b_h3, b_h3, K=K) - elif K > 64: - b_scale = _state_qdq_scale(b_h1, b_h2, b_h2, b_h2, K=K) - else: - b_scale = _state_qdq_scale(b_h1, b_h1, b_h1, b_h1, K=K) - b_h1 = fp8_scalar_qdq(b_h1, b_scale) - if K > 64: - b_h2 = fp8_scalar_qdq(b_h2, b_scale) - if K > 128: - b_h3 = fp8_scalar_qdq(b_h3, b_scale) - if K > 192: - b_h4 = fp8_scalar_qdq(b_h4, b_scale) - - # main recurrence - for i_t in range(NT): - i_t_int64 = i_t.to(tl.int64) - o_t = i_t_int64 * BT + tl.arange(0, BT) - m_t = o_t < T - if STATE_V_FIRST: - p_h1 = h + i_t_int64 * HV * K * V + o_v[:, None] * K + o_k1[None, :] - m_h1 = m_v[:, None] & m_k1[None, :] - else: - p_h1 = h + i_t_int64 * HV * K * V + o_k1[:, None] * V + o_v[None, :] - m_h1 = m_k1[:, None] & m_v[None, :] - tl.store(p_h1, b_h1.to(p_h1.dtype.element_ty), mask=m_h1) - if K > 64: - if STATE_V_FIRST: - p_h2 = h + i_t_int64 * HV * K * V + o_v[:, None] * K + o_k2[None, :] - m_h2 = m_v[:, None] & m_k2[None, :] - else: - p_h2 = h + i_t_int64 * HV * K * V + o_k2[:, None] * V + o_v[None, :] - m_h2 = m_k2[:, None] & m_v[None, :] - tl.store(p_h2, b_h2.to(p_h2.dtype.element_ty), mask=m_h2) - if K > 128: - if STATE_V_FIRST: - p_h3 = h + i_t_int64 * HV * K * V + o_v[:, None] * K + o_k3[None, :] - m_h3 = m_v[:, None] & m_k3[None, :] - else: - p_h3 = h + i_t_int64 * HV * K * V + o_k3[:, None] * V + o_v[None, :] - m_h3 = m_k3[:, None] & m_v[None, :] - tl.store(p_h3, b_h3.to(p_h3.dtype.element_ty), mask=m_h3) - if K > 192: - if STATE_V_FIRST: - p_h4 = h + i_t_int64 * HV * K * V + o_v[:, None] * K + o_k4[None, :] - m_h4 = m_v[:, None] & m_k4[None, :] - else: - p_h4 = h + i_t_int64 * HV * K * V + o_k4[:, None] * V + o_v[None, :] - m_h4 = m_k4[:, None] & m_v[None, :] - tl.store(p_h4, b_h4.to(p_h4.dtype.element_ty), mask=m_h4) - - p_w = w + o_t[:, None] * (HV * K) + o_k1[None, :] - b_w = tl.load(p_w, mask=m_t[:, None] & m_k1[None, :], other=0.0) - if STATE_V_FIRST: - b_v = tl.dot(b_w, tl.trans(b_h1).to(b_w.dtype)) - else: - b_v = tl.dot(b_w, b_h1.to(b_w.dtype)) - if K > 64: - p_w = w + o_t[:, None] * (HV * K) + o_k2[None, :] - b_w = tl.load(p_w, mask=m_t[:, None] & m_k2[None, :], other=0.0) - if STATE_V_FIRST: - b_v = tl.dot(b_w, tl.trans(b_h2).to(b_w.dtype), b_v) - else: - b_v = tl.dot(b_w, b_h2.to(b_w.dtype), b_v) - if K > 128: - p_w = w + o_t[:, None] * (HV * K) + o_k3[None, :] - b_w = tl.load(p_w, mask=m_t[:, None] & m_k3[None, :], other=0.0) - if STATE_V_FIRST: - b_v = tl.dot(b_w, tl.trans(b_h3).to(b_w.dtype), b_v) - else: - b_v = tl.dot(b_w, b_h3.to(b_w.dtype), b_v) - if K > 192: - p_w = w + o_t[:, None] * (HV * K) + o_k4[None, :] - b_w = tl.load(p_w, mask=m_t[:, None] & m_k4[None, :], other=0.0) - if STATE_V_FIRST: - b_v = tl.dot(b_w, tl.trans(b_h4).to(b_w.dtype), b_v) - else: - b_v = tl.dot(b_w, b_h4.to(b_w.dtype), b_v) - p_v = v + o_t[:, None] * (HV * V) + o_v[None, :] - b_v = tl.load(p_v, mask=m_t[:, None] & m_v[None, :], other=0.0) - b_v - - if SAVE_NEW_VALUE: - p_v = v_new + o_t[:, None] * (HV * V) + o_v[None, :] - tl.store(p_v, b_v.to(p_v.dtype.element_ty), mask=m_t[:, None] & m_v[None, :]) - - last_idx = min((i_t + 1) * BT, T) - 1 - if USE_G: - b_g_last = tl.load(g + (bos * HV + last_idx * HV + i_h).to(tl.int64)).to(tl.float32) - p_g = g + (bos * HV + i_h).to(tl.int64) + o_t * HV - b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32) - b_v = b_v * tl.where(m_t, exp2(b_g_last - b_g), 0)[:, None] - b_g_last = exp2(b_g_last) - b_h1 *= b_g_last - if K > 64: - b_h2 *= b_g_last - if K > 128: - b_h3 *= b_g_last - if K > 192: - b_h4 *= b_g_last - - if USE_GK: - o_k1 = tl.arange(0, 64) - b_gk_last1 = tl.load( - gk + (bos + last_idx) * HV * K + i_h * K + o_k1, mask=(o_k1 < K), other=0.0 - ).to(tl.float32) - if STATE_V_FIRST: - b_h1 *= exp2(b_gk_last1)[None, :] - else: - b_h1 *= exp2(b_gk_last1)[:, None] - if K > 64: - o_k2 = 64 + o_k1 - b_gk_last2 = tl.load( - gk + (bos + last_idx) * HV * K + i_h * K + o_k2, mask=(o_k2 < K), other=0.0 - ).to(tl.float32) - if STATE_V_FIRST: - b_h2 *= exp2(b_gk_last2)[None, :] - else: - b_h2 *= exp2(b_gk_last2)[:, None] - if K > 128: - o_k3 = 128 + o_k1 - b_gk_last3 = tl.load( - gk + (bos + last_idx) * HV * K + i_h * K + o_k3, mask=(o_k3 < K), other=0.0 - ).to(tl.float32) - if STATE_V_FIRST: - b_h3 *= exp2(b_gk_last3)[None, :] - else: - b_h3 *= exp2(b_gk_last3)[:, None] - if K > 192: - o_k4 = 192 + o_k1 - b_gk_last4 = tl.load( - gk + (bos + last_idx) * HV * K + i_h * K + o_k4, mask=(o_k4 < K), other=0.0 - ).to(tl.float32) - if STATE_V_FIRST: - b_h4 *= exp2(b_gk_last4)[None, :] - else: - b_h4 *= exp2(b_gk_last4)[:, None] - b_v = b_v.to(k.dtype.element_ty) - - p_k = k + o_k1[:, None] + o_t[None, :] * (H * K) - b_k = tl.load(p_k, mask=m_k1[:, None] & m_t[None, :], other=0.0) - if STATE_V_FIRST: - b_h1 += tl.trans(tl.dot(b_k, b_v)) - else: - b_h1 = tl.dot(b_k, b_v, b_h1) - if K > 64: - p_k = k + o_k2[:, None] + o_t[None, :] * (H * K) - b_k = tl.load(p_k, mask=m_k2[:, None] & m_t[None, :], other=0.0) - if STATE_V_FIRST: - b_h2 += tl.trans(tl.dot(b_k, b_v)) - else: - b_h2 = tl.dot(b_k, b_v, b_h2) - if K > 128: - p_k = k + o_k3[:, None] + o_t[None, :] * (H * K) - b_k = tl.load(p_k, mask=m_k3[:, None] & m_t[None, :], other=0.0) - if STATE_V_FIRST: - b_h3 += tl.trans(tl.dot(b_k, b_v)) - else: - b_h3 = tl.dot(b_k, b_v, b_h3) - if K > 192: - p_k = k + o_k4[:, None] + o_t[None, :] * (H * K) - b_k = tl.load(p_k, mask=m_k4[:, None] & m_t[None, :], other=0.0) - if STATE_V_FIRST: - b_h4 += tl.trans(tl.dot(b_k, b_v)) - else: - b_h4 = tl.dot(b_k, b_v, b_h4) - - # [ModelOpt] Dynamic per-tile FP8 QDQ of the next-chunk or stored final state. - # Recompute amax over [K, BV]; fp8_scalar_qdq applies that tile's scalar scale. - if STATE_QDQ == 1: - if K > 192: - b_scale = _state_qdq_scale(b_h1, b_h2, b_h3, b_h4, K=K) - elif K > 128: - b_scale = _state_qdq_scale(b_h1, b_h2, b_h3, b_h3, K=K) - elif K > 64: - b_scale = _state_qdq_scale(b_h1, b_h2, b_h2, b_h2, K=K) - else: - b_scale = _state_qdq_scale(b_h1, b_h1, b_h1, b_h1, K=K) - b_h1 = fp8_scalar_qdq(b_h1, b_scale) - if K > 64: - b_h2 = fp8_scalar_qdq(b_h2, b_scale) - if K > 128: - b_h3 = fp8_scalar_qdq(b_h3, b_scale) - if K > 192: - b_h4 = fp8_scalar_qdq(b_h4, b_scale) - - if STORE_FINAL_STATE: - if STATE_V_FIRST: - p_ht = ht + o_v[:, None] * K + o_k1[None, :] - m_ht = m_v[:, None] & m_k1[None, :] - else: - p_ht = ht + o_k1[:, None] * V + o_v[None, :] - m_ht = m_k1[:, None] & m_v[None, :] - tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), mask=m_ht) - if K > 64: - if STATE_V_FIRST: - p_ht = ht + o_v[:, None] * K + o_k2[None, :] - m_ht = m_v[:, None] & m_k2[None, :] - else: - p_ht = ht + o_k2[:, None] * V + o_v[None, :] - m_ht = m_k2[:, None] & m_v[None, :] - tl.store(p_ht, b_h2.to(p_ht.dtype.element_ty), mask=m_ht) - if K > 128: - if STATE_V_FIRST: - p_ht = ht + o_v[:, None] * K + o_k3[None, :] - m_ht = m_v[:, None] & m_k3[None, :] - else: - p_ht = ht + o_k3[:, None] * V + o_v[None, :] - m_ht = m_k3[:, None] & m_v[None, :] - tl.store(p_ht, b_h3.to(p_ht.dtype.element_ty), mask=m_ht) - if K > 192: - if STATE_V_FIRST: - p_ht = ht + o_v[:, None] * K + o_k4[None, :] - m_ht = m_v[:, None] & m_k4[None, :] - else: - p_ht = ht + o_k4[:, None] * V + o_v[None, :] - m_ht = m_k4[:, None] & m_v[None, :] - tl.store(p_ht, b_h4.to(p_ht.dtype.element_ty), mask=m_ht) - - -@triton.heuristics( - { - "USE_G": lambda args: args["g"] is not None, - "USE_GK": lambda args: args["gk"] is not None, - "USE_INITIAL_STATE": lambda args: args["dh0"] is not None, - "USE_FINAL_STATE_GRADIENT": lambda args: args["dht"] is not None, - "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, - } -) -@fla_cache_autotune( - configs=[ - triton.Config({"BV": BV}, num_warps=num_warps, num_stages=num_stages) - for num_warps in [2, 4] - for num_stages in ([2, 3, 4] if check_shared_mem("ampere") else [1]) - for BV in ([32, 64] if check_shared_mem("ada") else [32]) - ], - key=["H", "HV", "K", "V", "BT", "BV", "USE_G", "STATE_V_FIRST"], - **autotune_cache_kwargs, -) -@triton.jit(do_not_specialize=["T"]) -def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64( - q, - k, - w, - g, - gk, - dht, - dh0, - do, - dh, - dv, - dv2, - cu_seqlens, - chunk_offsets, - scale, - T, - H: tl.constexpr, - HV: tl.constexpr, - K: tl.constexpr, - V: tl.constexpr, - BT: tl.constexpr, - BV: tl.constexpr, - USE_G: tl.constexpr, - USE_GK: tl.constexpr, - USE_INITIAL_STATE: tl.constexpr, - USE_FINAL_STATE_GRADIENT: tl.constexpr, - STATE_V_FIRST: tl.constexpr, - IS_VARLEN: tl.constexpr, -): - pid = tl.program_id(0) - NV = tl.cdiv(V, BV) - i_v, i_nh = pid % NV, (pid // NV).to(tl.int64) - i_n, i_h = i_nh // HV, i_nh % HV - if IS_VARLEN: - bos, eos = ( - tl.load(cu_seqlens + i_n).to(tl.int64), - tl.load(cu_seqlens + i_n + 1).to(tl.int64), - ) - T = eos - bos - NT = tl.cdiv(T, BT) - boh = tl.load(chunk_offsets + i_n).to(tl.int64) - else: - bos, eos = i_n * T, i_n * T + T - NT = tl.cdiv(T, BT) - boh = i_n * NT - - if STATE_V_FIRST: - b_dh1 = tl.zeros([BV, 64], dtype=tl.float32) - if K > 64: - b_dh2 = tl.zeros([BV, 64], dtype=tl.float32) - if K > 128: - b_dh3 = tl.zeros([BV, 64], dtype=tl.float32) - if K > 192: - b_dh4 = tl.zeros([BV, 64], dtype=tl.float32) - else: - b_dh1 = tl.zeros([64, BV], dtype=tl.float32) - if K > 64: - b_dh2 = tl.zeros([64, BV], dtype=tl.float32) - if K > 128: - b_dh3 = tl.zeros([64, BV], dtype=tl.float32) - if K > 192: - b_dh4 = tl.zeros([64, BV], dtype=tl.float32) - - # calculate offset - q += (bos * H + i_h // (HV // H)).to(tl.int64) * K - k += (bos * H + i_h // (HV // H)).to(tl.int64) * K - w += (bos * HV + i_h).to(tl.int64) * K - do += (bos * HV + i_h).to(tl.int64) * V - dv += (bos * HV + i_h).to(tl.int64) * V - dv2 += (bos * HV + i_h).to(tl.int64) * V - dh += (boh * HV + i_h).to(tl.int64) * K * V - if USE_GK: - gk += (bos * HV + i_h).to(tl.int64) * K - - if USE_INITIAL_STATE: - dh0 += i_nh * K * V - if USE_FINAL_STATE_GRADIENT: - dht += i_nh * K * V - - o_v = i_v * BV + tl.arange(0, BV) - m_v = o_v < V - o_k1 = tl.arange(0, 64) - m_k1 = o_k1 < K - o_k2 = 64 + o_k1 - m_k2 = o_k2 < K - o_k3 = 128 + o_k1 - m_k3 = o_k3 < K - o_k4 = 192 + o_k1 - m_k4 = o_k4 < K - if USE_FINAL_STATE_GRADIENT: - if STATE_V_FIRST: - p_dht1 = dht + o_v[:, None] * K + o_k1[None, :] - m_dht1 = m_v[:, None] & m_k1[None, :] - else: - p_dht1 = dht + o_k1[:, None] * V + o_v[None, :] - m_dht1 = m_k1[:, None] & m_v[None, :] - b_dh1 += tl.load(p_dht1, mask=m_dht1, other=0.0) - if K > 64: - if STATE_V_FIRST: - p_dht2 = dht + o_v[:, None] * K + o_k2[None, :] - m_dht2 = m_v[:, None] & m_k2[None, :] - else: - p_dht2 = dht + o_k2[:, None] * V + o_v[None, :] - m_dht2 = m_k2[:, None] & m_v[None, :] - b_dh2 += tl.load(p_dht2, mask=m_dht2, other=0.0) - if K > 128: - if STATE_V_FIRST: - p_dht3 = dht + o_v[:, None] * K + o_k3[None, :] - m_dht3 = m_v[:, None] & m_k3[None, :] - else: - p_dht3 = dht + o_k3[:, None] * V + o_v[None, :] - m_dht3 = m_k3[:, None] & m_v[None, :] - b_dh3 += tl.load(p_dht3, mask=m_dht3, other=0.0) - if K > 192: - if STATE_V_FIRST: - p_dht4 = dht + o_v[:, None] * K + o_k4[None, :] - m_dht4 = m_v[:, None] & m_k4[None, :] - else: - p_dht4 = dht + o_k4[:, None] * V + o_v[None, :] - m_dht4 = m_k4[:, None] & m_v[None, :] - b_dh4 += tl.load(p_dht4, mask=m_dht4, other=0.0) - - for i_t in range(NT - 1, -1, -1): - i_t_int64 = i_t.to(tl.int64) - o_t = i_t_int64 * BT + tl.arange(0, BT) - m_t = o_t < T - if STATE_V_FIRST: - p_dh1 = dh + i_t_int64 * HV * K * V + o_v[:, None] * K + o_k1[None, :] - m_dh1 = m_v[:, None] & m_k1[None, :] - else: - p_dh1 = dh + i_t_int64 * HV * K * V + o_k1[:, None] * V + o_v[None, :] - m_dh1 = m_k1[:, None] & m_v[None, :] - tl.store(p_dh1, b_dh1.to(p_dh1.dtype.element_ty), mask=m_dh1) - if K > 64: - if STATE_V_FIRST: - p_dh2 = dh + i_t_int64 * HV * K * V + o_v[:, None] * K + o_k2[None, :] - m_dh2 = m_v[:, None] & m_k2[None, :] - else: - p_dh2 = dh + i_t_int64 * HV * K * V + o_k2[:, None] * V + o_v[None, :] - m_dh2 = m_k2[:, None] & m_v[None, :] - tl.store(p_dh2, b_dh2.to(p_dh2.dtype.element_ty), mask=m_dh2) - if K > 128: - if STATE_V_FIRST: - p_dh3 = dh + i_t_int64 * HV * K * V + o_v[:, None] * K + o_k3[None, :] - m_dh3 = m_v[:, None] & m_k3[None, :] - else: - p_dh3 = dh + i_t_int64 * HV * K * V + o_k3[:, None] * V + o_v[None, :] - m_dh3 = m_k3[:, None] & m_v[None, :] - tl.store(p_dh3, b_dh3.to(p_dh3.dtype.element_ty), mask=m_dh3) - if K > 192: - if STATE_V_FIRST: - p_dh4 = dh + i_t_int64 * HV * K * V + o_v[:, None] * K + o_k4[None, :] - m_dh4 = m_v[:, None] & m_k4[None, :] - else: - p_dh4 = dh + i_t_int64 * HV * K * V + o_k4[:, None] * V + o_v[None, :] - m_dh4 = m_k4[:, None] & m_v[None, :] - tl.store(p_dh4, b_dh4.to(p_dh4.dtype.element_ty), mask=m_dh4) - - last_idx = min((i_t_int64 + 1) * BT, T) - 1 - if USE_G: - bg_last = tl.load(g + (bos + last_idx) * HV + i_h).to(tl.float32) - p_g = g + bos * HV + i_h + o_t * HV - b_g = tl.load(p_g, mask=m_t, other=0.0).to(tl.float32) - bg_last_exp = exp2(bg_last) - b_g_exp = exp2(b_g) - p_dv = dv + o_t[:, None] * (HV * V) + o_v[None, :] - p_dv2 = dv2 + o_t[:, None] * (HV * V) + o_v[None, :] - p_do = do + o_t[:, None] * (HV * V) + o_v[None, :] - - b_do = tl.load(p_do, mask=m_t[:, None] & m_v[None, :], other=0.0) - - # Update dv - p_k = k + o_t[:, None] * (H * K) + o_k1[None, :] - b_k = tl.load(p_k, mask=m_t[:, None] & m_k1[None, :], other=0.0) - if USE_GK: - o_k1 = tl.arange(0, 64) - b_gk_last1 = tl.load(gk + last_idx * HV * K + o_k1, mask=(o_k1 < K), other=0.0).to( - tl.float32 - ) - if STATE_V_FIRST: - b_dv = tl.dot(b_k, tl.trans(b_dh1).to(b_k.dtype)) - else: - b_dv = tl.dot(b_k, b_dh1.to(b_k.dtype)) - - if K > 64: - p_k = k + o_t[:, None] * (H * K) + o_k2[None, :] - b_k = tl.load(p_k, mask=m_t[:, None] & m_k2[None, :], other=0.0) - if USE_GK: - b_gk_last2 = tl.load(gk + last_idx * HV * K + o_k2, mask=(o_k2 < K), other=0.0).to( - tl.float32 - ) - if STATE_V_FIRST: - b_dv = tl.dot(b_k, tl.trans(b_dh2).to(b_k.dtype), b_dv) - else: - b_dv = tl.dot(b_k, b_dh2.to(b_k.dtype), b_dv) - - if K > 128: - p_k = k + o_t[:, None] * (H * K) + o_k3[None, :] - b_k = tl.load(p_k, mask=m_t[:, None] & m_k3[None, :], other=0.0) - if USE_GK: - b_gk_last3 = tl.load(gk + last_idx * HV * K + o_k3, mask=(o_k3 < K), other=0.0).to( - tl.float32 - ) - if STATE_V_FIRST: - b_dv = tl.dot(b_k, tl.trans(b_dh3).to(b_k.dtype), b_dv) - else: - b_dv = tl.dot(b_k, b_dh3.to(b_k.dtype), b_dv) - - if K > 192: - p_k = k + o_t[:, None] * (H * K) + o_k4[None, :] - b_k = tl.load(p_k, mask=m_t[:, None] & m_k4[None, :], other=0.0) - if USE_GK: - b_gk_last4 = tl.load(gk + last_idx * HV * K + o_k4, mask=(o_k4 < K), other=0.0).to( - tl.float32 - ) - if STATE_V_FIRST: - b_dv = tl.dot(b_k, tl.trans(b_dh4).to(b_k.dtype), b_dv) - else: - b_dv = tl.dot(b_k, b_dh4.to(b_k.dtype), b_dv) - - if USE_G: - b_dv *= tl.where(m_t, exp2(bg_last - b_g), 0)[:, None] - b_dv += tl.load(p_dv, mask=m_t[:, None] & m_v[None, :], other=0.0) - - tl.store(p_dv2, b_dv.to(p_dv.dtype.element_ty), mask=m_t[:, None] & m_v[None, :]) - # Update dh - p_w = w + o_k1[:, None] + o_t[None, :] * (HV * K) - p_q = q + o_k1[:, None] + o_t[None, :] * (H * K) - b_w = tl.load(p_w, mask=m_k1[:, None] & m_t[None, :], other=0.0) - b_q = tl.load(p_q, mask=m_k1[:, None] & m_t[None, :], other=0.0) - if USE_G: - b_dh1 *= bg_last_exp - b_q = b_q * b_g_exp[None, :] - if USE_GK: - if STATE_V_FIRST: - b_dh1 *= exp2(b_gk_last1)[None, :] - else: - b_dh1 *= exp2(b_gk_last1[:, None]) - if STATE_V_FIRST: - b_dh1 += tl.trans( - tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - - tl.dot(b_w, b_dv.to(b_w.dtype)) - ) - else: - b_dh1 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot( - b_w, b_dv.to(b_w.dtype) - ) - if K > 64: - p_q = q + o_k2[:, None] + o_t[None, :] * (H * K) - p_w = w + o_k2[:, None] + o_t[None, :] * (HV * K) - b_q = tl.load(p_q, mask=m_k2[:, None] & m_t[None, :], other=0.0) - b_w = tl.load(p_w, mask=m_k2[:, None] & m_t[None, :], other=0.0) - if USE_G: - b_dh2 *= bg_last_exp - b_q = b_q * b_g_exp[None, :] - if USE_GK: - if STATE_V_FIRST: - b_dh2 *= exp2(b_gk_last2)[None, :] - else: - b_dh2 *= exp2(b_gk_last2[:, None]) - if STATE_V_FIRST: - b_dh2 += tl.trans( - tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - - tl.dot(b_w, b_dv.to(b_w.dtype)) - ) - else: - b_dh2 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot( - b_w, b_dv.to(b_w.dtype) - ) - if K > 128: - p_q = q + o_k3[:, None] + o_t[None, :] * (H * K) - p_w = w + o_k3[:, None] + o_t[None, :] * (HV * K) - b_q = tl.load(p_q, mask=m_k3[:, None] & m_t[None, :], other=0.0) - b_w = tl.load(p_w, mask=m_k3[:, None] & m_t[None, :], other=0.0) - if USE_G: - b_dh3 *= bg_last_exp - b_q = b_q * b_g_exp[None, :] - if USE_GK: - if STATE_V_FIRST: - b_dh3 *= exp2(b_gk_last3)[None, :] - else: - b_dh3 *= exp2(b_gk_last3[:, None]) - if STATE_V_FIRST: - b_dh3 += tl.trans( - tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - - tl.dot(b_w, b_dv.to(b_w.dtype)) - ) - else: - b_dh3 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot( - b_w, b_dv.to(b_w.dtype) - ) - if K > 192: - p_q = q + o_k4[:, None] + o_t[None, :] * (H * K) - p_w = w + o_k4[:, None] + o_t[None, :] * (HV * K) - b_q = tl.load(p_q, mask=m_k4[:, None] & m_t[None, :], other=0.0) - b_w = tl.load(p_w, mask=m_k4[:, None] & m_t[None, :], other=0.0) - if USE_G: - b_dh4 *= bg_last_exp - b_q = b_q * b_g_exp[None, :] - if USE_GK: - if STATE_V_FIRST: - b_dh4 *= exp2(b_gk_last4)[None, :] - else: - b_dh4 *= exp2(b_gk_last4[:, None]) - if STATE_V_FIRST: - b_dh4 += tl.trans( - tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - - tl.dot(b_w, b_dv.to(b_w.dtype)) - ) - else: - b_dh4 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot( - b_w, b_dv.to(b_w.dtype) - ) - - if USE_INITIAL_STATE: - if STATE_V_FIRST: - p_dh0 = dh0 + o_v[:, None] * K + o_k1[None, :] - m_dh0 = m_v[:, None] & m_k1[None, :] - else: - p_dh0 = dh0 + o_k1[:, None] * V + o_v[None, :] - m_dh0 = m_k1[:, None] & m_v[None, :] - tl.store(p_dh0, b_dh1.to(p_dh0.dtype.element_ty), mask=m_dh0) - if K > 64: - if STATE_V_FIRST: - p_dh1 = dh0 + o_v[:, None] * K + o_k2[None, :] - m_dh1 = m_v[:, None] & m_k2[None, :] - else: - p_dh1 = dh0 + o_k2[:, None] * V + o_v[None, :] - m_dh1 = m_k2[:, None] & m_v[None, :] - tl.store(p_dh1, b_dh2.to(p_dh1.dtype.element_ty), mask=m_dh1) - if K > 128: - if STATE_V_FIRST: - p_dh2 = dh0 + o_v[:, None] * K + o_k3[None, :] - m_dh2 = m_v[:, None] & m_k3[None, :] - else: - p_dh2 = dh0 + o_k3[:, None] * V + o_v[None, :] - m_dh2 = m_k3[:, None] & m_v[None, :] - tl.store(p_dh2, b_dh3.to(p_dh2.dtype.element_ty), mask=m_dh2) - if K > 192: - if STATE_V_FIRST: - p_dh3 = dh0 + o_v[:, None] * K + o_k4[None, :] - m_dh3 = m_v[:, None] & m_k4[None, :] - else: - p_dh3 = dh0 + o_k4[:, None] * V + o_v[None, :] - m_dh3 = m_k4[:, None] & m_v[None, :] - tl.store(p_dh3, b_dh4.to(p_dh3.dtype.element_ty), mask=m_dh3) - - -def chunk_gated_delta_rule_fwd_h( - k: torch.Tensor, - w: torch.Tensor, - u: torch.Tensor, - g: torch.Tensor | None = None, - gk: torch.Tensor | None = None, - initial_state: torch.Tensor | None = None, - output_final_state: bool = False, - chunk_size: int = 64, - save_new_value: bool = True, - state_v_first: bool = False, - cu_seqlens: torch.LongTensor | None = None, - cu_seqlens_cpu: torch.LongTensor | None = None, - chunk_indices: torch.LongTensor | None = None, - chunk_offsets: torch.LongTensor | None = None, - state_qdq: int = STATE_QDQ_OFF, - state_qdq_block_v: int | None = None, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: - B, T, H, K, V, HV = *k.shape, u.shape[-1], u.shape[2] - BT = chunk_size - BV = state_qdq_tile_v(V, state_qdq, state_qdq_block_v) - - if chunk_indices is None and cu_seqlens is not None: - chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) - # N: the actual number of sequences in the batch with either equal or variable lengths - if cu_seqlens is None: - N, NT, chunk_offsets = B, triton.cdiv(T, BT), None - else: - N, NT = len(cu_seqlens) - 1, len(chunk_indices) - if chunk_offsets is None: - chunk_offsets = prepare_chunk_offsets(cu_seqlens, BT) - assert K <= 256, "current kernel does not support head dimension larger than 256." - - if state_v_first: - h = k.new_empty(B, NT, HV, V, K) - final_state = k.new_zeros(N, HV, V, K, dtype=torch.float32) if output_final_state else None - else: - h = k.new_empty(B, NT, HV, K, V) - final_state = k.new_zeros(N, HV, K, V, dtype=torch.float32) if output_final_state else None - - v_new = torch.empty_like(u) if save_new_value else None - - grid = (triton.cdiv(V, BV) * N * HV,) - chunk_gated_delta_rule_fwd_kernel_h_blockdim64[grid]( - k=k, - v=u, - w=w, - v_new=v_new, - g=g, - gk=gk, - h=h, - h0=initial_state, - ht=final_state, - cu_seqlens=cu_seqlens, - chunk_offsets=chunk_offsets, - T=T, - H=H, - HV=HV, - K=K, - V=V, - BT=BT, - BV=BV, - STATE_V_FIRST=state_v_first, - STATE_QDQ=state_qdq, - ) - return h, v_new, final_state - - -def state_qdq_tile_v(V: int, state_qdq: int, state_qdq_block_v: int | None) -> int: - """Return the V tile width ``BV`` of the forward state kernel. - - Without state quantization this is fla's largest tile. With it, the tile is also the - quantization granularity: one dynamic scale per ``[K, BV]`` block of a head's state. The - default is fla's 64-column tile, i.e. one scale per sequence and head for ``V <= 64`` and two - for the usual ``V == 128``. ``state_qdq_block_v=128`` gives one scale per 128-wide head but - exceeds the register budget where the kernel is limited to two warps (Blackwell) and spills. - """ - if state_qdq == STATE_QDQ_OFF: - return 64 if check_shared_mem("ada") else 32 - if state_qdq != STATE_QDQ_FP8_DYNAMIC: - raise ValueError(f"Unsupported state_qdq mode {state_qdq}; expected 0 or 1.") - BV = min(triton.next_power_of_2(V), 64) if state_qdq_block_v is None else state_qdq_block_v - if BV < 16 or BV > STATE_QDQ_MAX_BLOCK_V or BV & (BV - 1): - raise ValueError( - f"state_qdq_block_v must be a power of two in [16, {STATE_QDQ_MAX_BLOCK_V}], got {BV}." - ) - return BV - - -def chunk_gated_delta_rule_bwd_dhu( - q: torch.Tensor, - k: torch.Tensor, - w: torch.Tensor, - do: torch.Tensor, - dv: torch.Tensor, - g: torch.Tensor | None = None, - gk: torch.Tensor | None = None, - h0: torch.Tensor | None = None, - dht: torch.Tensor | None = None, - scale: float | None = None, - state_v_first: bool = False, - cu_seqlens: torch.LongTensor | None = None, - chunk_size: int = 64, - chunk_indices: torch.LongTensor | None = None, - chunk_offsets: torch.LongTensor | None = None, - use_graph: bool = False, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - B, T, H, K, V, HV = *q.shape, do.shape[-1], do.shape[2] - # N: the actual number of sequences in the batch with either equal or variable lengths - BT = chunk_size - assert K <= 256, "current kernel does not support head dimension being larger than 256." - - if chunk_indices is None and cu_seqlens is not None: - chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) - if cu_seqlens is None: - N, NT, chunk_offsets = B, triton.cdiv(T, BT), None - else: - N, NT = len(cu_seqlens) - 1, len(chunk_indices) - if chunk_offsets is None: - chunk_offsets = prepare_chunk_offsets(cu_seqlens, BT) - - if use_graph: - # [ModelOpt] Imported here: fla.ops.utils.graph is absent from released fla versions. - from fla.ops.utils.graph import get_static_buffer - - if state_v_first: - dh = get_static_buffer("dhu_dh_vf", (B, NT, HV, V, K), q.dtype, q.device) - else: - dh = get_static_buffer("dhu_dh", (B, NT, HV, K, V), q.dtype, q.device) - dh0 = ( - get_static_buffer("dhu_dh0", tuple(h0.shape), torch.float32, h0.device) - if h0 is not None - else None - ) - dv2 = get_static_buffer("dhu_dv2", tuple(dv.shape), dv.dtype, dv.device) - else: - if state_v_first: - dh = q.new_empty(B, NT, HV, V, K) - else: - dh = q.new_empty(B, NT, HV, K, V) - dh0 = torch.empty_like(h0, dtype=torch.float32) if h0 is not None else None - dv2 = torch.empty_like(dv) - - def grid(meta): - return (triton.cdiv(V, meta["BV"]) * N * HV,) - - chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64[grid]( - q=q, - k=k, - w=w, - g=g, - gk=gk, - dht=dht, - dh0=dh0, - do=do, - dh=dh, - dv=dv, - dv2=dv2, - cu_seqlens=cu_seqlens, - chunk_offsets=chunk_offsets, - scale=scale, - T=T, - H=H, - HV=HV, - K=K, - V=V, - BT=BT, - STATE_V_FIRST=state_v_first, - ) - return dh, dh0, dv2 diff --git a/modelopt/torch/kernels/quantization/linear_attention/fla_chunk_gated_delta_rule.py b/modelopt/torch/kernels/quantization/linear_attention/fla_chunk_gated_delta_rule.py deleted file mode 100644 index a5609cab102..00000000000 --- a/modelopt/torch/kernels/quantization/linear_attention/fla_chunk_gated_delta_rule.py +++ /dev/null @@ -1,721 +0,0 @@ -# Adapted from: https://github.com/fla-org/flash-linear-attention/blob/516143e31fce/fla/ops/gated_delta_rule/chunk.py -# Adapted with modifications (marked [ModelOpt]): threads state_qdq / state_qdq_block_v through -# the autograd function, applies an optional w_quantizer to the WY tensor w, and imports the -# state kernels from the vendored sibling module. -# -# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li -# -# This source code is licensed under the MIT license found in the -# LICENSE file in the root directory of this source tree. -# For a list of all contributors, visit: -# https://github.com/fla-org/flash-linear-attention/graphs/contributors - - -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 AND MIT -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import warnings - -import fla -import torch -from fla.modules.l2norm import l2norm_bwd, l2norm_fwd -from fla.ops.common.chunk_o import chunk_bwd_dqkwg, chunk_bwd_dv_local, chunk_fwd_o -from fla.ops.common.gate import fused_beta_sigmoid, fused_beta_sigmoid_bwd -from fla.ops.cp import FLACPContext -from fla.ops.cp.chunk_delta_h import ( - chunk_gated_delta_rule_bwd_dhu_pre_process, - chunk_gated_delta_rule_fwd_h_pre_process, - compress_h0, - expand_h0, -) -from fla.ops.gated_delta_rule.chunk_fwd import chunk_gated_delta_rule_fwd_intra -from fla.ops.gated_delta_rule.gate import gdn_gate_bwd, gdn_gate_chunk_cumsum -from fla.ops.gated_delta_rule.wy_fast import prepare_wy_repr_bwd, recompute_w_u_fwd -from fla.ops.utils import chunk_local_cumsum -from fla.ops.utils.constant import RCP_LN2 -from fla.ops.utils.index import prepare_chunk_indices -from fla.utils import ( - IS_NVIDIA_HOPPER, - TRITON_ABOVE_3_4_0, - autocast_custom_bwd, - autocast_custom_fwd, - input_guard, -) - -from modelopt.torch.quantization.linear_attention.utils import validate_gdn_quantizer -from modelopt.torch.quantization.nn import TensorQuantizer - -from .fla_chunk_delta_h import ( - STATE_QDQ_FP8_DYNAMIC, - STATE_QDQ_OFF, - chunk_gated_delta_rule_bwd_dhu, - chunk_gated_delta_rule_fwd_h, -) - - -def chunk_gated_delta_rule_fwd( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - g: torch.Tensor, - beta: torch.Tensor, - scale: float, - initial_state: torch.Tensor, - output_final_state: bool, - state_v_first: bool = False, - cu_seqlens: torch.LongTensor | None = None, - cp_context: FLACPContext | None = None, - chunk_indices: torch.LongTensor | None = None, - use_gate_in_kernel: bool = False, - A_log: torch.Tensor | None = None, - dt_bias: torch.Tensor | None = None, - chunk_size: int = 64, - state_qdq: int = STATE_QDQ_OFF, - state_qdq_block_v: int | None = None, - w_quantizer: TensorQuantizer | None = None, -): - g_input = g if use_gate_in_kernel else None - if use_gate_in_kernel: - g = gdn_gate_chunk_cumsum( - g=g, - A_log=A_log, - chunk_size=chunk_size, - scale=RCP_LN2, - dt_bias=dt_bias, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - ) - else: - g = chunk_local_cumsum( - g, - chunk_size=chunk_size, - scale=RCP_LN2, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - ) - # obtain WY representation. u is actually the new v. - # fused kkt + solve_tril + recompute_w_u - w, u, A = chunk_gated_delta_rule_fwd_intra( - k=k, - v=v, - g=g, - beta=beta, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - chunk_size=chunk_size, - ) - # [ModelOpt] w is an activation (the WY form of the chunk's keys) that lands in memory here, - # so it is fake-quantized once per forward instead of tile by tile inside the kernel. - if w_quantizer is not None: - w = w_quantizer(w) - - if cp_context is not None: - initial_state = chunk_gated_delta_rule_fwd_h_pre_process( - k=k, - w=w, - u=u, - g=g, - cu_seqlens=cu_seqlens, - initial_state=initial_state, - context=cp_context, - state_v_first=state_v_first, - chunk_size=chunk_size, - ) - - h, v_new, final_state = chunk_gated_delta_rule_fwd_h( - k=k, - w=w, - u=u, - g=g, - initial_state=initial_state, - output_final_state=output_final_state, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - state_v_first=state_v_first, - chunk_size=chunk_size, - state_qdq=state_qdq, - state_qdq_block_v=state_qdq_block_v, - ) - - if cp_context is not None: - initial_state = compress_h0(initial_state, context=cp_context) - - o = chunk_fwd_o( - q=q, - k=k, - v=v_new, - h=h, - g=g, - scale=scale, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - state_v_first=state_v_first, - chunk_size=chunk_size, - ) - return g, o, A, final_state, initial_state, g_input, w if w_quantizer is not None else None - - -def chunk_gated_delta_rule_bwd( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - g: torch.Tensor, - beta: torch.Tensor, - A: torch.Tensor, - scale: float, - initial_state: torch.Tensor, - do: torch.Tensor, - dht: torch.Tensor, - state_v_first: bool = False, - cu_seqlens: torch.LongTensor | None = None, - cp_context: FLACPContext | None = None, - chunk_indices: torch.LongTensor | None = None, - use_gate_in_kernel: bool = False, - g_input: torch.Tensor | None = None, - A_log: torch.Tensor | None = None, - dt_bias: torch.Tensor | None = None, - chunk_size: int = 64, - state_qdq: int = STATE_QDQ_OFF, - state_qdq_block_v: int | None = None, - quantized_w: torch.Tensor | None = None, -): - w, u = recompute_w_u_fwd( - k=k, - v=v, - beta=beta, - A=A, - g=g, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - ) - # [ModelOpt] Reuse forward QDQ exactly; the validated policy uses identity STE. - # This also avoids calling observers or dynamic scale computation in backward. - if quantized_w is not None: - w = quantized_w - - if cp_context is not None: - initial_state = expand_h0(initial_state, context=cp_context) - - # [ModelOpt] The backward recomputes the forward's chunk states, so it sees the same - # fake-quantized states; the state gradient itself passes straight through the QDQ. - h, v_new, _ = chunk_gated_delta_rule_fwd_h( - k=k, - w=w, - u=u, - g=g, - initial_state=initial_state, - output_final_state=False, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - state_v_first=state_v_first, - chunk_size=chunk_size, - state_qdq=state_qdq, - state_qdq_block_v=state_qdq_block_v, - ) - dv = chunk_bwd_dv_local( - q=q, - k=k, - g=g, - do=do, - scale=scale, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - chunk_size=chunk_size, - ) - - if cp_context is not None: - # initial_state is None in the CP mode - # We only need to compute dht of current rank and pass it to the backward kernel - dht, initial_state = chunk_gated_delta_rule_bwd_dhu_pre_process( - q=q, - k=k, - w=w, - do=do, - dv=dv, - g=g, - scale=scale, - cu_seqlens=cu_seqlens, - dht=dht, - initial_state=initial_state, - context=cp_context, - state_v_first=state_v_first, - chunk_size=chunk_size, - ) - - dh, dh0, dv = chunk_gated_delta_rule_bwd_dhu( - q=q, - k=k, - w=w, - g=g, - h0=initial_state, - dht=dht, - do=do, - dv=dv, - scale=scale, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - state_v_first=state_v_first, - chunk_size=chunk_size, - ) - dq, dk, dw, dg = chunk_bwd_dqkwg( - q=q, - k=k, - v=v_new, - w=w, - g=g, - h=h, - dv=dv, - do=do, - dh=dh, - scale=scale, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - state_v_first=state_v_first, - chunk_size=chunk_size, - ) - dk2, dv, db, dg2 = prepare_wy_repr_bwd( - k=k, - v=v, - beta=beta, - g=g, - A=A, - dw=dw, - du=dv, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - ) - dk.add_(dk2) - dg.add_(dg2) - dg = chunk_local_cumsum( - dg, chunk_size=chunk_size, reverse=True, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices - ) - dA_log, ddt_bias = None, None - if use_gate_in_kernel: - dg, dA_log, ddt_bias = gdn_gate_bwd(g=g_input, A_log=A_log, dt_bias=dt_bias, dyg=dg) - return dq, dk, dv, db, dg, dh0, dA_log, ddt_bias - - -class ChunkGatedDeltaRuleFunction(torch.autograd.Function): - @staticmethod - @input_guard - @autocast_custom_fwd - def forward( - ctx, - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - g: torch.Tensor, - beta: torch.Tensor, - scale: float, - initial_state: torch.Tensor, - output_final_state: bool, - state_v_first: bool = False, - cu_seqlens: torch.LongTensor | None = None, - cu_seqlens_cpu: torch.LongTensor | None = None, - chunk_indices: torch.LongTensor | None = None, - use_qk_l2norm_in_kernel: bool = False, - use_gate_in_kernel: bool = False, - A_log: torch.Tensor | None = None, - dt_bias: torch.Tensor | None = None, - use_beta_sigmoid_in_kernel: bool = False, - allow_neg_eigval: bool = False, - cp_context: FLACPContext | None = None, - chunk_size: int = 64, - state_qdq: int = STATE_QDQ_OFF, - state_qdq_block_v: int | None = None, - w_quantizer: TensorQuantizer | None = None, - ): - q_rstd, k_rstd = None, None - if use_qk_l2norm_in_kernel: - q, q_rstd = l2norm_fwd(q) - k, k_rstd = l2norm_fwd(k) - - beta_raw = beta - if use_beta_sigmoid_in_kernel: - beta = fused_beta_sigmoid(beta_raw, scale=2.0 if allow_neg_eigval else 1.0) - - if chunk_indices is None and cu_seqlens is not None: - chunk_indices = prepare_chunk_indices( - cu_seqlens, chunk_size, cu_seqlens_cpu=cu_seqlens_cpu - ) - g, o, A, final_state, initial_state, g_input, quantized_w = chunk_gated_delta_rule_fwd( - q=q, - k=k, - v=v, - g=g, - beta=beta, - scale=scale, - initial_state=initial_state, - output_final_state=output_final_state, - cu_seqlens=cu_seqlens, - cp_context=cp_context, - chunk_indices=chunk_indices, - state_v_first=state_v_first, - use_gate_in_kernel=use_gate_in_kernel, - A_log=A_log, - dt_bias=dt_bias, - chunk_size=chunk_size, - state_qdq=state_qdq, - state_qdq_block_v=state_qdq_block_v, - w_quantizer=w_quantizer, - ) - ctx.save_for_backward( - q, - q_rstd, - k, - k_rstd, - v, - g, - beta_raw, - beta, - A, - initial_state, - cu_seqlens, - chunk_indices, - g_input, - A_log, - dt_bias, - quantized_w, - ) - ctx.scale = scale - ctx.chunk_size = chunk_size - ctx.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel - ctx.use_beta_sigmoid_in_kernel = use_beta_sigmoid_in_kernel - ctx.allow_neg_eigval = allow_neg_eigval - ctx.cp_context = cp_context - ctx.state_v_first = state_v_first - ctx.use_gate_in_kernel = use_gate_in_kernel - ctx.state_qdq = state_qdq - ctx.state_qdq_block_v = state_qdq_block_v - return o.to(q.dtype), final_state - - @staticmethod - @input_guard - @autocast_custom_bwd - def backward( - ctx, - do: torch.Tensor, - dht: torch.Tensor, - ): - ( - q, - q_rstd, - k, - k_rstd, - v, - g, - beta_raw, - beta, - A, - initial_state, - cu_seqlens, - chunk_indices, - g_input, - A_log, - dt_bias, - quantized_w, - ) = ctx.saved_tensors - dq, dk, dv, db, dg, dh0, dA_log, ddt_bias = chunk_gated_delta_rule_bwd( - q=q, - k=k, - v=v, - g=g, - beta=beta, - A=A, - scale=ctx.scale, - initial_state=initial_state, - do=do, - dht=dht, - cu_seqlens=cu_seqlens, - cp_context=ctx.cp_context, - chunk_indices=chunk_indices, - state_v_first=ctx.state_v_first, - use_gate_in_kernel=ctx.use_gate_in_kernel, - g_input=g_input, - A_log=A_log, - dt_bias=dt_bias, - chunk_size=ctx.chunk_size, - state_qdq=ctx.state_qdq, - state_qdq_block_v=ctx.state_qdq_block_v, - quantized_w=quantized_w, - ) - if ctx.use_qk_l2norm_in_kernel: - dq = l2norm_bwd(q, q_rstd, dq) - dk = l2norm_bwd(k, k_rstd, dk) - if ctx.use_beta_sigmoid_in_kernel: - db = fused_beta_sigmoid_bwd(beta_raw, db, scale=2.0 if ctx.allow_neg_eigval else 1.0) - return ( - dq.to(q), - dk.to(k), - dv.to(v), - dg.to(g), - db.to(beta_raw), - None, - dh0, - None, - None, - None, - None, - None, - None, - None, - dA_log, - ddt_bias, - None, - None, - None, - None, - None, - None, - None, - ) - - -# [ModelOpt] Not registered with fla's backend dispatch: another backend must not take over a -# call that asks for state quantization. -@torch.compiler.disable -def chunk_gated_delta_rule( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - g: torch.Tensor, - beta: torch.Tensor, - scale: float | None = None, - initial_state: torch.Tensor | None = None, - output_final_state: bool = False, - use_qk_l2norm_in_kernel: bool = False, - use_beta_sigmoid_in_kernel: bool = False, - allow_neg_eigval: bool = False, - state_v_first: bool = False, - cu_seqlens: torch.LongTensor | None = None, - cu_seqlens_cpu: torch.LongTensor | None = None, - chunk_indices: torch.LongTensor | None = None, - cp_context: FLACPContext | None = None, - **kwargs, -): - r""" - Args: - q (torch.Tensor): - queries of shape `[B, T, H, K]`. - k (torch.Tensor): - keys of shape `[B, T, H, K]`. - v (torch.Tensor): - values of shape `[B, T, HV, V]`. - GVA (Grouped Value Attention) is applied if `HV > H`, where `HV` must be divisible by `H`. - g (torch.Tensor): - (forget) gating tensor of shape `[B, T, HV]`. - When `use_gate_in_kernel=False` (default), `g` should be in log space (pre-computed decay). - When `use_gate_in_kernel=True`, `g` is the raw input before gate activation; - the kernel fuses `-exp(A_log) * softplus(g + dt_bias)` + chunk cumsum internally. - beta (torch.Tensor): - betas of shape `[B, T, HV]`. - scale (Optional[float]): - Scale factor for the RetNet attention scores. - If not provided, it will default to `1 / sqrt(K)`. Default: `None`. - initial_state (Optional[torch.Tensor]): - Initial state of shape `[N, HV, K, V]` for `N` input sequences. - For equal-length input sequences, `N` equals the batch size `B`. - Default: `None`. - output_final_state (Optional[bool]): - Whether to output the final state of shape `[N, HV, K, V]`. Default: `False`. - use_qk_l2norm_in_kernel (bool): - Whether to apply L2norm to the q/k tensor internally. Default: `False`. - use_gate_in_kernel (bool): - Whether to compute the log-space GDN decay internally. - When `True`, the passed `g` is the raw input, and `A_log` must be provided. - The kernel fuses gate activation + chunk cumsum in a single pass. - Default: `False`. - A_log (Optional[torch.Tensor]): - Decay parameter of shape `[HV]`. Required when `use_gate_in_kernel=True`. - dt_bias (Optional[torch.Tensor]): - Bias added to `g` before activation, of shape `[HV]`. - Only used when `use_gate_in_kernel=True`. - use_beta_sigmoid_in_kernel (bool): - Whether to apply `torch.sigmoid(beta)` before launching the chunk kernel. - - If `True`, the passed `beta` acts as the raw beta logits. - - If `False`, `beta` is expected to already be in post-sigmoid space. - Default: `False`. - allow_neg_eigval (bool): - Whether to allow negative eigenvalues by scaling `beta` to `[0, 2)`. - Only takes effect together with `use_beta_sigmoid_in_kernel=True`, in which case - the kernel computes `2 * sigmoid(beta)` instead of `sigmoid(beta)`. Default: `False`. - state_v_first (Optional[bool]): - Store the recurrent state in V-first ``[V, K]`` layout instead of the default ``[K, V]``. Default: ``False``. - cu_seqlens (torch.LongTensor): - Cumulative sequence lengths of shape `[N+1]` used for variable-length training, - consistent with the FlashAttention API. - chunk_indices (Optional[torch.LongTensor]): - Pre-computed chunk indices for variable-length inputs. - If provided, they are used directly instead of being computed from `cu_seqlens`. Default: `None`. - cp_context (Optional[FLACPContext]): - Context parallel context for distributed training across multiple devices. - When provided, `initial_state` and `output_final_state` are not supported, - and `cu_seqlens` will be overridden by the context. Default: `None`. - - Returns: - o (torch.Tensor): - Outputs of shape `[B, T, HV, V]`. - final_state (torch.Tensor): - Final state of shape `[N, HV, K, V]` if `output_final_state=True` else `None`. - - Examples:: - >>> import torch - >>> import torch.nn.functional as F - >>> from einops import rearrange - >>> from fla.ops.gated_delta_rule import chunk_gated_delta_rule - # inputs with equal lengths - >>> B, T, H, HV, K, V = 4, 2048, 4, 8, 512, 512 - >>> q = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda') - >>> k = F.normalize(torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda'), p=2, dim=-1) - >>> v = torch.randn(B, T, HV, V, dtype=torch.bfloat16, device='cuda') - >>> beta = torch.rand(B, T, HV, dtype=torch.bfloat16, device='cuda').sigmoid() - >>> g = F.logsigmoid(torch.rand(B, T, HV, dtype=torch.bfloat16, device='cuda')) - >>> h0 = torch.randn(B, HV, K, V, dtype=torch.bfloat16, device='cuda') - >>> o, ht = chunk_gated_delta_rule( - q, k, v, g, beta, - initial_state=h0, - output_final_state=True - ) - # for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required - >>> q, k, v, beta, g = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, beta, g)) - # for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected - >>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long) - >>> o, ht = chunk_gated_delta_rule( - q, k, v, g, beta, - initial_state=h0, - output_final_state=True, - cu_seqlens=cu_seqlens - ) - """ - if fla.__version__ != "0.5.1": - raise RuntimeError(f"ModelOpt GDN requires fla-core==0.5.1, got {fla.__version__}.") - if "transpose_state_layout" in kwargs: - if state_v_first: - raise ValueError( - "Cannot pass both `state_v_first` and the deprecated `transpose_state_layout`." - ) - warnings.warn( - "`transpose_state_layout` is deprecated and renamed to `state_v_first`.", - DeprecationWarning, - stacklevel=2, - ) - state_v_first = kwargs.pop("transpose_state_layout") - - # Validate head dimensions - if q.shape[2] != k.shape[2]: - raise ValueError( - f"q and k must have the same number of heads, " - f"but got q.shape[2]={q.shape[2]} and k.shape[2]={k.shape[2]}" - ) - H, HV = q.shape[2], v.shape[2] - if HV % H != 0: - raise ValueError( - f"For GVA, num_v_heads (HV={HV}) must be evenly divisible by " - f"num_heads (H={H}), but got HV % H = {HV % H}" - ) - - if "head_first" in kwargs: - raise DeprecationWarning( - "head_first has been removed. Inputs must be in `[B, T, H, ...]` format.", - ) - - chunk_size = kwargs.pop("chunk_size", 64) - if chunk_size != 64: - raise ValueError("ModelOpt GDN supports only chunk_size=64; FLA WY backward assumes 64.") - - # [ModelOpt] state_qdq: 0 keeps fla's numerics; 1 fake-quantizes the state carried between - # chunks to FP8 E4M3 with a dynamic scale per [K, state_qdq_block_v] tile of each head. - state_qdq = kwargs.pop("state_qdq", STATE_QDQ_OFF) - state_qdq_block_v = kwargs.pop("state_qdq_block_v", None) - # w_quantizer: dynamic FP8 TensorQuantizer applied to the WY tensor - # ``w`` of shape [B, T, HV, K] before it multiplies the state, emulating an FP8 x FP8 matmul. - w_quantizer = kwargs.pop("w_quantizer", None) - use_gate_in_kernel = kwargs.pop("use_gate_in_kernel", False) - A_log = kwargs.pop("A_log", None) - dt_bias = kwargs.pop("dt_bias", None) - if kwargs: - raise TypeError(f"Unexpected keyword arguments: {', '.join(sorted(kwargs))}") - if state_qdq not in (STATE_QDQ_OFF, STATE_QDQ_FP8_DYNAMIC): - raise ValueError(f"`state_qdq` must be 0 or 1, got {state_qdq}.") - if w_quantizer is not None: - validate_gdn_quantizer(w_quantizer, name="gdn_w_quantizer") - if state_qdq and (not q.is_cuda or torch.cuda.get_device_capability(q.device) < (8, 9)): - raise RuntimeError("GDN state QDQ requires native E4M3 conversion on CUDA SM89 or newer.") - if (state_qdq != STATE_QDQ_OFF or w_quantizer is not None) and cp_context is not None: - raise ValueError("State or w quantization is not supported together with `cp_context`.") - - if cp_context is not None: - assert initial_state is None, "Initial state is not supported for CP" - assert output_final_state is False, "Output final state is not supported for CP" - assert cp_context.cu_seqlens is not None, "cu_seqlens is required for CP" - cu_seqlens = cp_context.cu_seqlens - if cp_context.cu_seqlens_cpu is not None: - cu_seqlens_cpu = cp_context.cu_seqlens_cpu - - if cu_seqlens is not None: - if q.shape[0] != 1: - raise ValueError( - f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`." - f"Please flatten variable-length inputs before processing.", - ) - if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1: - raise ValueError( - f"The number of initial states is expected to be equal to the number of input sequences, " - f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}.", - ) - if use_gate_in_kernel: - assert A_log is not None, "A_log must be provided when use_gate_in_kernel=True." - if allow_neg_eigval and not use_beta_sigmoid_in_kernel: - raise ValueError("`allow_neg_eigval=True` requires `use_beta_sigmoid_in_kernel=True`.") - - if scale is None: - scale = k.shape[-1] ** -0.5 - # [ModelOpt] Hopper's TileLang backward needs BF16 and equal head counts. Expand outside - # custom autograd so repeat_interleave reduces q/k gradients back to the original heads. - if IS_NVIDIA_HOPPER and TRITON_ABOVE_3_4_0: - if any(x.dtype != torch.bfloat16 for x in (q, k, v)): - raise ValueError("Hopper with Triton >= 3.4 requires BF16 q/k/v for GDN training.") - if H != HV: - q = q.repeat_interleave(HV // H, dim=2) - k = k.repeat_interleave(HV // H, dim=2) - o, final_state = ChunkGatedDeltaRuleFunction.apply( - q, - k, - v, - g, - beta, - scale, - initial_state, - output_final_state, - state_v_first, - cu_seqlens, - cu_seqlens_cpu, - chunk_indices, - use_qk_l2norm_in_kernel, - use_gate_in_kernel, - A_log, - dt_bias, - use_beta_sigmoid_in_kernel, - allow_neg_eigval, - cp_context, - chunk_size, - state_qdq, - state_qdq_block_v, - w_quantizer, - ) - return o, final_state - - -chunk_gdn = chunk_gated_delta_rule diff --git a/modelopt/torch/kernels/quantization/linear_attention/serving/__init__.py b/modelopt/torch/kernels/quantization/linear_attention/serving/__init__.py new file mode 100644 index 00000000000..1156582d24e --- /dev/null +++ b/modelopt/torch/kernels/quantization/linear_attention/serving/__init__.py @@ -0,0 +1,25 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Optional vLLM forward kernels for linear-attention serving precision profiles.""" + +try: + import vllm +except ImportError as error: + raise ImportError( + "Linear-attention serving profiles require the optional vLLM dependency: " + "public vLLM for 'vllm' (legacy spelling 'vllm_0_15'), or the compatible quantized-ReplaySSM " + "vLLM fork for 'replayssm'." + ) from error diff --git a/modelopt/torch/kernels/quantization/linear_attention/serving/_compat.py b/modelopt/torch/kernels/quantization/linear_attention/serving/_compat.py new file mode 100644 index 00000000000..04153111bdb --- /dev/null +++ b/modelopt/torch/kernels/quantization/linear_attention/serving/_compat.py @@ -0,0 +1,46 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Keep native vLLM imports and state layout changes at the kernel boundary.""" + +from functools import cache +from importlib import import_module +from importlib.util import find_spec + +import vllm + +__all__ = [] + +_FLA_OPS = ( + "vllm.third_party.flash_linear_attention.ops" + if find_spec("vllm.third_party.flash_linear_attention") is not None + else "vllm.model_executor.layers.fla.ops" +) + + +@cache +def state_v_first(): + """Inspect the native layout at first use, after optional-dependency docs imports.""" + # vLLM 0.16 transposed the recurrent cache; ModelOpt keeps [H,K,V] for QDQ/autograd. + return vllm.__version_tuple__[:2] >= (0, 16) + + +def fla_module(name): + return import_module(f"{_FLA_OPS}.{name}") + + +def state_layout(state): + """Convert between ModelOpt and native layout; transposing is its own inverse.""" + return state.transpose(-1, -2).contiguous() if state_v_first() else state.contiguous() diff --git a/modelopt/torch/kernels/quantization/linear_attention/serving/chunk_delta_h.py b/modelopt/torch/kernels/quantization/linear_attention/serving/chunk_delta_h.py new file mode 100644 index 00000000000..08e43dc6489 --- /dev/null +++ b/modelopt/torch/kernels/quantization/linear_attention/serving/chunk_delta_h.py @@ -0,0 +1,77 @@ +# Adapted from: https://github.com/vllm-project/vllm/blob/1892993bc18e243e2c05841314c5e9c06a80c70d/vllm/model_executor/layers/fla/ops/chunk_delta_h.py +# Modifications: import the kernel; retain FP32 intermediates for the training adjoint. +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-License-Identifier: Apache-2.0 +# Original FLA code: Copyright (c) 2023-2025, Songlin Yang, Yu Zhang. +# FLA is licensed under the MIT license reproduced in the root LICENSE. + +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 AND MIT +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Save FP32 training intermediates using vLLM's unchanged chunk-state kernel.""" + +from functools import cache +from inspect import signature + +import torch +import triton + +from ._compat import fla_module, state_v_first + +_chunk = fla_module("chunk_delta_h") +chunk_gated_delta_rule_fwd_kernel_h_blockdim64 = ( + _chunk.chunk_gated_delta_rule_fwd_kernel_h_blockdim64 +) +prepare_chunk_offsets = fla_module("index").prepare_chunk_offsets + + +@cache +def _has_exp2(): + """Inspect the native signature only when a kernel runs, not during docs imports.""" + return "use_exp2" in signature(_chunk.chunk_gated_delta_rule_fwd_h).parameters + + +def chunk_state(k, w, u, g, gk, initial_state, cu_seqlens, use_exp2=False): + """Return chunk-start states, residual values, and final state for one sequence.""" + _, length, heads, key_dim = k.shape + value_dim = u.shape[-1] + # The Torch adjoint needs unrounded values. Output kernels receive BF16 casts. + state_shape = (value_dim, key_dim) if state_v_first() else (key_dim, value_dim) + h = k.new_empty(1, triton.cdiv(length, 64), heads, *state_shape, dtype=torch.float32) + updated = torch.empty_like(u, dtype=torch.float32) + final = torch.empty_like(initial_state, dtype=torch.float32) + chunk_gated_delta_rule_fwd_kernel_h_blockdim64[ + lambda meta: (triton.cdiv(value_dim, meta["BV"]), heads) + ]( + k=k, + v=u, + w=w, + v_new=updated, + g=g, + gk=gk, + h=h, + h0=initial_state, + ht=final, + cu_seqlens=cu_seqlens, + chunk_offsets=prepare_chunk_offsets(cu_seqlens, 64), + T=length, + H=heads, + Hg=heads, + K=key_dim, + V=value_dim, + BT=64, + **({"USE_EXP2": use_exp2} if _has_exp2() else {}), + ) + return h, updated, final diff --git a/modelopt/torch/kernels/quantization/linear_attention/serving/forward.py b/modelopt/torch/kernels/quantization/linear_attention/serving/forward.py new file mode 100644 index 00000000000..1d0866ca21f --- /dev/null +++ b/modelopt/torch/kernels/quantization/linear_attention/serving/forward.py @@ -0,0 +1,146 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""vLLM forward primitives with FP32 saved values; autograd lives in quantization.""" + +from functools import cache +from importlib import import_module + +import torch + +from ._compat import fla_module, state_layout +from .chunk_delta_h import chunk_state + +chunk_fwd_o = fla_module("chunk_o").chunk_fwd_o +chunk_scaled_dot_kkt_fwd = fla_module("chunk_scaled_dot_kkt").chunk_scaled_dot_kkt_fwd +chunk_local_cumsum = fla_module("cumsum").chunk_local_cumsum +fused_recurrent_gated_delta_rule = fla_module("fused_recurrent").fused_recurrent_gated_delta_rule +_kda = fla_module("kda") +chunk_gla_fwd_o_gk = _kda.chunk_gla_fwd_o_gk +chunk_kda_scaled_dot_kkt_fwd = _kda.chunk_kda_scaled_dot_kkt_fwd +fused_recurrent_kda = _kda.fused_recurrent_kda +fused_kda_gate = _kda.fused_kda_gate +kda_wu = _kda.recompute_w_u_fwd +l2norm_fwd = fla_module("l2norm").l2norm_fwd +solve_tril = fla_module("solve_tril").solve_tril +gdn_wu = fla_module("wy_fast").recompute_w_u_fwd +_KDA_GATE_SCALE = getattr(_kda, "RCP_LN2", 1.0) + + +def prefill(q, k, v, g, beta, state, scale, normalize=False): + """Return one sequence's output, final state, and rounded forward intermediates.""" + q, k, v, g, beta = [x.unsqueeze(0).contiguous() for x in (q, k, v, g, beta)] + cu = torch.tensor([0, q.shape[1]], device=q.device, dtype=torch.int32) + if normalize: + q, k = l2norm_fwd(q), l2norm_fwd(k) + channel = g.ndim == 4 + gc = chunk_local_cumsum(g, chunk_size=64, cu_seqlens=cu) + # Newer native KDA kernels evaluate exp2 of base-2 cumulative log gates. + natural_gc = gc + if channel: + gc = gc * _KDA_GATE_SCALE + if channel: + lower, scores = chunk_kda_scaled_dot_kkt_fwd(q, k, gc, beta, scale=scale, cu_seqlens=cu) + else: + lower = chunk_scaled_dot_kkt_fwd( + k=k, beta=beta, g=gc, cu_seqlens=cu, output_dtype=torch.float32 + ) + scores = None + inverse = solve_tril(A=lower, cu_seqlens=cu, output_dtype=k.dtype) + if channel: + w, u, _, kg = kda_wu(k=k, v=v, beta=beta, A=inverse, gk=gc, cu_seqlens=cu) + else: + w, u = gdn_wu(k=k, v=v, beta=beta, A=inverse, g_cumsum=gc, cu_seqlens=cu) + kg = k + assert kg is not None + h, updated, final = chunk_state( + k=kg, + w=w, + u=u, + g=None if channel else gc, + gk=gc if channel else None, + initial_state=state_layout(state.unsqueeze(0)), + cu_seqlens=cu, + use_exp2=channel and _KDA_GATE_SCALE != 1.0, + ) + assert updated is not None and final is not None + if channel: + out = chunk_gla_fwd_o_gk( + q=q, + v=updated.to(v.dtype), + g=gc, + A=scores, + h=h.to(k.dtype), + scale=scale, + o=torch.empty_like(v), + cu_seqlens=cu, + chunk_size=64, + ) + else: + out = chunk_fwd_o( + q=q, k=k, v=updated.to(v.dtype), h=h.to(k.dtype), g=gc, scale=scale, cu_seqlens=cu + ) + intermediates = { + "q": q[0], + "k": k[0], + "g": natural_gc[0], + "lower": lower[0], + "inverse": inverse[0], + "w": w[0], + "u": u[0], + "h": state_layout(h[0]), + "updated": updated[0], + "kg": kg[0], + "scores": None if scores is None else scores[0], + } + return out[0], state_layout(final[0]), intermediates + + +def step(q, k, v, g, beta, state, scale, normalize=False): + """Run one native token update without modifying the incoming state.""" + recurrent = fused_recurrent_kda if g.ndim == 2 else fused_recurrent_gated_delta_rule + out, final = recurrent( + *[x[None, None].contiguous() for x in (q, k, v, g, beta)], + initial_state=state_layout(state.unsqueeze(0)), + inplace_final_state=False, + scale=scale, + use_qk_l2norm_in_kernel=normalize, + cu_seqlens=torch.tensor([0, 1], device=q.device, dtype=torch.int32), + # Private dense training state needs no paged-cache index (0 is now reserved). + ssm_state_indices=None, + ) + return out[0, 0], state_layout(final[0]) + + +def fused_gdn_gating(*args, **kwargs): + """Load the optional vLLM model only when Megatron needs GDN gate preparation.""" + return _gdn_gate()(*args, **kwargs) + + +# Megatron compiles gate preparation; module discovery must execute outside that graph. +@torch.compiler.disable +@cache +def _gdn_gate(): + for path in ( + "vllm.model_executor.layers.mamba.gdn.qwen_gdn_linear_attn", + "vllm.model_executor.layers.mamba.gdn_linear_attn", + "vllm.model_executor.models.qwen3_next", + ): + try: + return import_module(path).fused_gdn_gating + except ModuleNotFoundError as error: # noqa: PERF203 - cached, one-time import discovery + if error.name is None or not (path == error.name or path.startswith(error.name + ".")): + raise + raise ImportError("The installed vLLM does not provide fused_gdn_gating") diff --git a/modelopt/torch/kernels/quantization/linear_attention/serving/replay.py b/modelopt/torch/kernels/quantization/linear_attention/serving/replay.py new file mode 100644 index 00000000000..17cd816a758 --- /dev/null +++ b/modelopt/torch/kernels/quantization/linear_attention/serving/replay.py @@ -0,0 +1,130 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Functional training launches of the optional quantized-ReplaySSM serving kernels.""" + +import torch +from vllm.model_executor.layers.fla.ops.fused_recurrent_replayssm import ( + _apply_h_value_axis, + fused_recurrent_gated_delta_rule_replayssm, + prefill_write_checkpoint, +) + + +def _storage(value, enabled): + heads, keys, values = value.shape + state = torch.zeros( + 2, heads, values, keys, device=value.device, dtype=torch.int8 if enabled else torch.float32 + ) + scales = ( + torch.zeros(2, heads, values // 32, keys, device=value.device, dtype=torch.float16) + if enabled + else None + ) + indices = torch.ones(1, device=value.device, dtype=torch.int32) + return state, scales, indices + + +def _decoded(state, scales): + values = state[1].transpose(-1, -2).float() + if scales is None: + return values, None + metadata = scales[1].transpose(-1, -2).contiguous() + return (values.unflatten(-1, (-1, 32)) * metadata.float().unsqueeze(-1)).flatten(-2), metadata + + +def checkpoint(value, enabled): + """Return the serving checkpoint's decoded Hadamard-basis values and FP16 scales.""" + state, scales, indices = _storage(value, enabled) + original = value.transpose(-1, -2).unsqueeze(0).contiguous().float() + if enabled: + prefill_write_checkpoint( + original, state, None, None, scales, None, None, indices, 8, hadamard=True + ) + else: + state[1] = _apply_h_value_axis(original)[0] + return _decoded(state, scales) + + +def original_basis(value): + """Use serving's FP32 checkpoint transform when returning a dense state.""" + return _apply_h_value_axis(value.transpose(-1, -2).unsqueeze(0).contiguous())[0].transpose( + -1, -2 + ) + + +def step(q, k, v, gate, beta, carry, window, enabled, scale, normalize, gate_inputs): + """Launch one token on private cache buffers; never mutate tensors saved for backward.""" + anchor = carry.anchor + state, scales, indices = _storage(anchor.values, enabled) + if enabled: + metadata = anchor.scales + codes = (anchor.values.unflatten(-1, (-1, 32)) / metadata.float().unsqueeze(-1)).round() + state[1] = codes.flatten(-2).transpose(-1, -2).to(torch.int8) + scales[1] = metadata.transpose(-1, -2) + else: + state[1] = anchor.values.transpose(-1, -2) + heads, keys, values = anchor.values.shape + d = torch.zeros(2, heads, window, values, device=q.device, dtype=torch.bfloat16) + kc = torch.zeros(2, heads, window, keys, device=q.device, dtype=torch.bfloat16) + channel = gate.ndim == 2 + gc = torch.zeros( + (2, heads, window, keys) if channel else (2, heads, window), + device=q.device, + dtype=torch.float32, + ) + for i, entry in enumerate(carry.entries): + d[1, :, i] = entry.update.values + kc[1, :, i] = entry.key.values + gc[1, :, i] = entry.log_retention + mixed = torch.cat([x.to(torch.bfloat16).flatten() for x in (q, k, v)])[None] + if channel: + a, b, a_log, bias = gate.flatten()[None], beta[None], None, None + else: + raw_gate, raw_beta, a_log, bias = gate_inputs + a, b = raw_gate[None].contiguous(), raw_beta[None].contiguous() + output = torch.empty(1, heads, values, device=q.device, dtype=torch.bfloat16) + cursor = carry.cursor + fused_recurrent_gated_delta_rule_replayssm( + mixed, + a, + b, + a_log, + bias, + scale, + state, + d, + kc, + gc, + output, + indices, + torch.full((1,), cursor, device=q.device, dtype=torch.int32), + use_qk_l2norm_in_kernel=normalize, + quant_state_bits=8 if enabled else 0, + state_scale=scales, + hadamard_value_basis=True, + block_v=64, + nk=2 if keys == 32 else None, + _is_kda=channel, + ) + decoded, metadata = _decoded(state, scales) + return ( + output[0].float(), + decoded, + metadata, + kc[1, :, cursor].float(), + d[1, :, cursor].float(), + gc[1, :, cursor], + ) diff --git a/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py b/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py index a4796162f6d..f1bd15154fb 100644 --- a/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py +++ b/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py @@ -60,6 +60,7 @@ def remove_per_module_state( _ = metadata.pop("subnet_config", None) _ = metadata.pop("real_quantizer_state", None) _ = metadata.pop("q_tensor_state", None) + _ = metadata.pop("linear_attention", None) else: config["metadata"] = {} diff --git a/modelopt/torch/quantization/config.py b/modelopt/torch/quantization/config.py index f8969ecfdc2..0e84b0e6bc9 100644 --- a/modelopt/torch/quantization/config.py +++ b/modelopt/torch/quantization/config.py @@ -161,6 +161,8 @@ from modelopt.torch.opt.config_loader import load_config from modelopt.torch.utils.network import ConstructorLike +from .linear_attention.config import LinearAttentionPolicyEntry + class QuantizerCfgEntry(ModeloptBaseConfig): """A single entry in a ``quant_cfg`` list.""" @@ -1568,6 +1570,11 @@ def _dict_to_entry(key: str, value) -> list[dict[str, Any]]: class QuantizeConfig(ModeloptBaseConfig): """Default configuration for ``quantize`` mode.""" + linear_attention: list[LinearAttentionPolicyEntry] = ModeloptField( + default=[], + title="Execution policies for supported linear-attention training modules", + ) + quant_cfg: QuantizeQuantCfgType = ModeloptField( default=[{"quantizer_name": "*", "cfg": {"num_bits": 8, "axis": None}}], title="Quantization configuration", diff --git a/modelopt/torch/quantization/conversion.py b/modelopt/torch/quantization/conversion.py index 78a6e03a16b..405f9b8876b 100644 --- a/modelopt/torch/quantization/conversion.py +++ b/modelopt/torch/quantization/conversion.py @@ -62,12 +62,16 @@ def convert_to_quantized_model(model: ModelLikeModule, config: QuantizeConfig) -> ConvertReturnType: """Convert the model to a quantized one as per `config`.""" + # Defer shared plugin imports to avoid the plugin/conversion initialization cycle. + from .plugins.linear_attention import _apply_linear_attention_policy, _validate_linear_attention + # initialize the true module if necessary model = model.init_modellike() if isinstance(model, ModelLikeModule) else model replace_quant_module(model, version=ModeloptStateManager(model).state_version) set_quantizer_by_cfg(model, config.get("quant_cfg", [])) - _validate_linear_attention_quantizers(model) + _apply_linear_attention_policy(model, config) + _validate_linear_attention(model) metadata = {} update_quantize_metadata(model, config, metadata) @@ -95,8 +99,15 @@ def restore_quantized_model( model: ModelLikeModule, config: QuantizeConfig, metadata: MetadataDict ) -> nn.Module: """Insert quantizers to the model and restore the quantizer states from the given state dict.""" + # Defer the shared plugin import to avoid the plugin/conversion initialization cycle. + from .plugins.linear_attention import _apply_linear_attention_policy + # initialize the true module if necessary - convert_to_quantized_model(model, config) + model = model.init_modellike() if isinstance(model, ModelLikeModule) else model + replace_quant_module(model, version=ModeloptStateManager(model).state_version) + set_quantizer_by_cfg(model, config.get("quant_cfg", [])) + # Rebuild from the recipe without validating against partially restored state. + _apply_linear_attention_policy(model, config) return restore_quantizer_state(model, config, metadata) @@ -137,6 +148,13 @@ def restore_quantizer_state(model: nn.Module, config: QuantizeConfig, metadata: details regarding how MCore sharded checkpoint is restored, see modelopt.torch.opt.plugins.mcore_dist_checkpointing.restore_sharded_modelopt_state. """ + # Defer shared plugin imports to avoid the plugin/conversion initialization cycle. + from .plugins.linear_attention import ( + _restore_legacy_linear_attention_quantizers, + _restore_linear_attention_policy, + ) + + _restore_linear_attention_policy(model, metadata.get("linear_attention")) if "quantizer_state" not in metadata: # MCore sharded checkpoint (`torch-dist`) has its quantizer_state stored as the # extra_state of `QuantModule`. The quantizer_state is resumed with @@ -144,13 +162,8 @@ def restore_quantizer_state(model: nn.Module, config: QuantizeConfig, metadata: return model quantizer_state_dict = dict(metadata["quantizer_state"]) - # Older checkpoints predate these disabled handles; preserve their baseline path. - for name, module in _linear_attention_modules(model).items(): - for handle in ("gdn_state_quantizer", "gdn_w_quantizer"): - key = f"{name}.{handle}" if name else handle - quantizer = getattr(module, handle) - if key not in quantizer_state_dict and not quantizer.is_enabled: - quantizer_state_dict[key] = quantizer.get_modelopt_state() + if "linear_attention" not in metadata: + _restore_legacy_linear_attention_quantizers(model, quantizer_state_dict) unmatched_keys = quantizer_state_dict.keys() - quantizer_state(model).keys() extra_keys = quantizer_state(model).keys() - quantizer_state_dict.keys() @@ -203,6 +216,14 @@ def update_quantize_metadata( model: nn.Module, config: QuantizeConfig, metadata: MetadataDict ) -> None: """Update the quantizer state in the metadata dict.""" + # Defer the shared plugin import to avoid the plugin/conversion initialization cycle. + from .plugins.linear_attention import _linear_attention_state + + policies = _linear_attention_state(model) + if policies: + metadata["linear_attention"] = policies + else: + metadata.pop("linear_attention", None) metadata["quantizer_state"] = quantizer_state(model) if shared_state_metadata := SharedWeightGlobalAmaxState.metadata(model): metadata["shared_quant_states"] = shared_state_metadata @@ -210,22 +231,6 @@ def update_quantize_metadata( metadata.pop("shared_quant_states", None) -def _linear_attention_modules(model): - # Optional framework plugins import conversion; defer this import to avoid that cycle. - from .plugins.gated_delta_net import GatedDeltaNetStateQuantMixin - - return { - get_unwrapped_name(name, model): module - for name, module in model.named_modules() - if isinstance(module, GatedDeltaNetStateQuantMixin) - } - - -def _validate_linear_attention_quantizers(model): - for module in _linear_attention_modules(model).values(): - module.validate_linear_attention() - - def quantizer_state(model: nn.Module) -> dict[str, Any]: """Returns the quantizer state dict describing the quantizer states in the model.""" return { diff --git a/modelopt/torch/quantization/linear_attention/__init__.py b/modelopt/torch/quantization/linear_attention/__init__.py index 477eeb89847..42bda38f4e4 100644 --- a/modelopt/torch/quantization/linear_attention/__init__.py +++ b/modelopt/torch/quantization/linear_attention/__init__.py @@ -13,4 +13,10 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Helpers for linear-attention quantization.""" +"""Numerical policies and differentiable kernels for linear attention.""" + +from .config import * +from .decode import * +from .gdn import * +from .kda import * +from .training import * diff --git a/modelopt/torch/quantization/linear_attention/_vllm_autograd.py b/modelopt/torch/quantization/linear_attention/_vllm_autograd.py new file mode 100644 index 00000000000..e25971171b5 --- /dev/null +++ b/modelopt/torch/quantization/linear_attention/_vllm_autograd.py @@ -0,0 +1,145 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Differentiable adapter for the installed serving runtime's arithmetic. + +The kernels supply forward values. Autograd differentiates the corresponding +operations evaluated at those values; operand casts and QDQ use identity STE. +The optional vLLM dependency supplies kernels; training owns its state and needs no server. +""" + +import torch + +from .utils import forward_value + + +def rounded(value): + return forward_value(value, value.to(torch.bfloat16)) + + +def normalized(value): + return value / (value.square().sum(-1, keepdim=True) + 1e-6).sqrt() + + +def prefix(q, k, v, g, beta, state, scale, beta_dtype, normalize=False): + # Keep the optional vLLM dependency isolated to the native precision profile. + from ...kernels.quantization.linear_attention.serving.forward import prefill + + with torch.no_grad(): + out, final, saved = prefill( + q.to(torch.bfloat16), + k.to(torch.bfloat16), + v.to(torch.bfloat16), + g.float(), + beta.to(beta_dtype), + state.float(), + scale, + normalize, + ) + if not torch.is_grad_enabled() or not any(x.requires_grad for x in (q, k, v, g, beta, state)): + return out.float(), final + q, k = [ + forward_value(normalized(x) if normalize else x, saved[n]) for n, x in (("q", q), ("k", k)) + ] + outputs = [] + channel = g.ndim == 3 + for chunk, lo in enumerate(range(0, len(q), 64)): + hi = min(lo + 64, len(q)) + qc, kc, vc = [x[lo:hi].transpose(0, 1) for x in (q, k, v)] + bc = beta[lo:hi].transpose(0, 1).unsqueeze(-1) + gc = forward_value(g[lo:hi].transpose(0, 1).cumsum(1), saved["g"][lo:hi].transpose(0, 1)) + state = forward_value(state, saved["h"][chunk]) + hs = rounded(state) + count = hi - lo + + def matrix(name, value): + return forward_value(value, saved[name][lo:hi].transpose(0, 1)[..., :count]) + + if channel: + lower_rows, score_rows = [], [] + for row in range(count): + right = ( + kc[:, : row + 1] * (gc[:, row : row + 1] - gc[:, : row + 1]).exp() + ).transpose(-1, -2) + lower_rows.append( + torch.nn.functional.pad( + (bc[:, row : row + 1] * kc[:, row : row + 1]) @ right, (0, count - row - 1) + ) + ) + score_rows.append( + torch.nn.functional.pad( + (qc[:, row : row + 1] * scale) @ right, (0, count - row - 1) + ) + ) + lower = matrix("lower", torch.cat(lower_rows, dim=1).tril(-1)) + scores = matrix("scores", torch.cat(score_rows, dim=1)) + gate = gc.exp() + else: + causal = torch.ones(count, count, device=q.device, dtype=torch.bool).tril() + decay = (gc.unsqueeze(-1) - gc.unsqueeze(-2)).masked_fill(~causal, 0).exp() + lower = matrix("lower", (bc * (kc @ kc.transpose(-1, -2)) * decay).tril(-1)) + scores = ((qc @ kc.transpose(-1, -2)) * decay).tril() + gate = gc.exp().unsqueeze(-1) + eye = torch.eye(count, device=q.device, dtype=q.dtype).expand_as(lower) + inverse = matrix( + "inverse", + torch.linalg.solve_triangular(eye + lower, eye, upper=False, unitriangular=True), + ) + u = forward_value(inverse @ rounded(bc * vc), saved["u"][lo:hi].transpose(0, 1)) + kb = rounded(bc * kc) if beta_dtype == torch.bfloat16 else bc * kc + w = forward_value(inverse @ rounded(kb * gate), saved["w"][lo:hi].transpose(0, 1)) + updated = forward_value(u - w @ hs, saved["updated"][lo:hi].transpose(0, 1)) + if channel: + output = rounded(rounded(qc * scale) * gate) @ hs + rounded(scores) @ rounded(updated) + kg = forward_value(kc * (gc[:, -1:] - gc).exp(), saved["kg"][lo:hi].transpose(0, 1)) + state = state * gate[:, -1, :, None] + kg.transpose(-1, -2) @ rounded(updated) + else: + output = ((qc @ hs) * gate + rounded(scores) @ rounded(updated)) * scale + weighted = rounded(updated * (gc[:, -1:] - gc).exp().unsqueeze(-1)) + state = state * gate[:, -1, :, None] + kc.transpose(-1, -2) @ weighted + outputs.append(forward_value(output.transpose(0, 1), out[lo:hi])) + return torch.cat(outputs), forward_value(state, final) + + +def step(q, k, v, gate, beta, state, scale, normalize=False): + # Keep the optional vLLM dependency isolated to the native precision profile. + from ...kernels.quantization.linear_attention.serving.forward import step as native_step + + with torch.no_grad(): + out, final = native_step( + q.to(torch.bfloat16), + k.to(torch.bfloat16), + v.to(torch.bfloat16), + gate.float(), + beta.float(), + state.float(), + scale, + normalize, + ) + if not torch.is_grad_enabled() or not any( + x.requires_grad for x in (q, k, v, gate, beta, state) + ): + return out.float(), final + if normalize: + q, k = normalized(q), normalized(k) + decay = gate.exp().unsqueeze(-1) + if gate.ndim == 1: + decay = decay.unsqueeze(-1) + decayed = state * decay + residual = v - (decayed * k.unsqueeze(-1)).sum(-2) + update = residual * beta.unsqueeze(-1) + working = forward_value(decayed + k.unsqueeze(-1) * update.unsqueeze(-2), final) + output = ((q * scale).unsqueeze(-1) * working).sum(-2) + return forward_value(output, out), working diff --git a/modelopt/torch/quantization/linear_attention/config.py b/modelopt/torch/quantization/linear_attention/config.py new file mode 100644 index 00000000000..a616a94dba3 --- /dev/null +++ b/modelopt/torch/quantization/linear_attention/config.py @@ -0,0 +1,72 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Saved execution policy for GDN/KDA state QAT.""" + +from typing import Literal + +from pydantic import Field, model_validator + +from modelopt.torch.opt.config import ModeloptBaseConfig, ModeloptField + +__all__ = ["LinearAttentionConfig", "LinearAttentionPolicyEntry"] + + +class LinearAttentionConfig(ModeloptBaseConfig): + """One execution policy shared by the chunked prefix and recurrent suffix. + + ``serving`` uses native forward arithmetic with a differentiable adjoint; + ``fla`` is the disabled/default path. ``vllm`` uses the installed vLLM's arithmetic; + ``vllm_0_15`` remains an accepted legacy spelling for existing checkpoints. + ``replayssm`` uses the serving fork's INT8/Hadamard checkpoint and ring kernels. + ``replay_window=1`` refreshes the state every token; larger windows enable replay. + Fresh prefill has no internal state QDQ; incoming continuation state and + the handoff to a nonempty suffix use the selected checkpoint encoding. + + TensorQuantizer independently enables QDQ and owns its format and grouping. + ``state_block_v`` controls grouping only for legacy per-tile quantizers. + """ + + backend: Literal["fla", "serving"] = ModeloptField(default="fla") + precision: Literal["vllm", "vllm_0_15", "replayssm"] = ModeloptField(default="vllm") + replay_window: int = Field(default=1, ge=1, le=64, strict=True) + state_block_v: Literal[16, 32, 64, 128] = ModeloptField(default=64) + + @property + def state_codec(self): + """The native profile fixes the codec; it is not an independent setting.""" + return "int8_hadamard32" if self.precision == "replayssm" else "tile" + + @model_validator(mode="after") + def _validate_profile(self): + if self.precision == "replayssm": + if self.backend != "serving": + raise ValueError("replayssm requires backend='serving'") + if self.state_block_v < 32: + raise ValueError("int8_hadamard32 requires state_block_v >= 32") + elif self.replay_window != 1: + raise ValueError("Replay requires precision='replayssm'") + return self + + +class LinearAttentionPolicyEntry(ModeloptBaseConfig): + """Assign a complete policy to supported modules matching ``module_name``. + + Rules apply in order: the last match wins, without merging fields. + A rule must match at least one supported module across distributed stages. + """ + + module_name: str = Field(...) + cfg: LinearAttentionConfig = ModeloptField(default=LinearAttentionConfig()) diff --git a/modelopt/torch/quantization/linear_attention/decode.py b/modelopt/torch/quantization/linear_attention/decode.py new file mode 100644 index 00000000000..622a825b900 --- /dev/null +++ b/modelopt/torch/quantization/linear_attention/decode.py @@ -0,0 +1,467 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Differentiable token-state and encoded-update replay for QAT.""" + +from dataclasses import dataclass + +import torch + +from .config import LinearAttentionConfig +from .utils import _resolve_state_quantizer, _tile_qdq, forward_value + +__all__ = [ + "EncodedLinearAttentionTensor", + "LinearAttentionCarry", + "LinearAttentionState", + "ReplayEntry", + "recurrent_decode", +] + + +@dataclass +class EncodedLinearAttentionTensor: + """Floating QDQ values with autograd history and optional quantization scales. + + ``format`` and ``block_v`` describe the emulated encoding; ``values`` stores + decoded values, not packed INT8/FP8 codes. Scale metadata is non-differentiable. + """ + + values: torch.Tensor + scales: torch.Tensor | None + format: str + block_v: int | None + + +@dataclass +class ReplayEntry: + """One saved rank-one update since the last state checkpoint. + + ``key`` is [H,K] and ``update`` is the beta-scaled residual [H,V]. ReplaySSM + stores their BF16 forward values once. ``log_retention`` is [H] for GDN or + [H,K] for KDA and records this token's decay for subsequent reconstruction. + """ + + key: EncodedLinearAttentionTensor + update: EncodedLinearAttentionTensor + log_retention: torch.Tensor + + +@dataclass +class LinearAttentionState: + """Runtime state passed between recurrent tokens and successive decode calls. + + With ``replay_window=1``, every token refreshes ``anchor`` and ``entries`` + stays empty. With a larger window, the anchor stays fixed while updates + accumulate; reaching the window refreshes the anchor and clears the entries. + Tensor values retain their autograd history for QAT/QAD across call boundaries. + """ + + anchor: EncodedLinearAttentionTensor # Checkpoint [H,K,V]; QDQ is applied when enabled. + entries: tuple[ReplayEntry, ...] # Updates since the checkpoint, in token order. + position: int # Next token position, including the supplied prefix offset. + started: bool # Initial-state encoding has run; resuming must not repeat it. + signature: str # Execution/quantizer settings used to check compatible reuse. + value_basis: str = "identity" # Anchor and updates may use the Hadamard value basis. + precision: str = "vllm" # Selects the native profile's reconstruction arithmetic. + + @property + def cursor(self): + """Number of encoded updates since the last anchor refresh.""" + return len(self.entries) + + def reconstruct(self, *, original_basis=True): + """Return the current dense [H,K,V] state without refreshing or quantizing it. + + Apply pending updates and decay to the anchor, preserving gradients. + Decode uses ``original_basis=False`` to continue in the stored basis; + the default converts Hadamard state back for the caller's final-state output. + Keep this object, rather than only this tensor, to resume a replay window. + """ + state = _replay_state(self) if self.precision == "replayssm" else self.anchor.values + if original_basis and self.value_basis == "hadamard32": + # The optional serving fork owns the checkpoint basis transform. + from ...kernels.quantization.linear_attention.serving.replay import ( + original_basis as restore_basis, + ) + + state = restore_basis(state) + return state + + +# Compatibility name for existing imports and pickled runtime state objects. +LinearAttentionCarry = LinearAttentionState + + +def _hadamard32(value): + """Apply the orthonormal Sylvester transform to contiguous 32-value groups.""" + shape = value.shape + for width in (1, 2, 4, 8, 16): + pairs = value.reshape(*shape[:-1], -1, 2, width) + left, right = pairs.unbind(-2) + value = torch.stack((left + right, left - right), dim=-2).reshape(shape) + return value * (32**-0.5) + + +def _replay_state(carry: LinearAttentionState, current_gate=None): + """Decay the anchor and add saved updates weighted by their subsequent gates. + + KDA weights keys per channel; GDN weights updates per head. ``current_gate`` + includes the next KDA token's decay before BF16 weighting, without adding its update. + """ + state = carry.anchor.values + if not carry.entries: + if current_gate is not None: + state = state * current_gate.exp().unsqueeze(-1) + return state + gates = torch.stack([entry.log_retention for entry in carry.entries], dim=1) + keys = torch.stack([entry.key.values for entry in carry.entries], dim=1) + updates = torch.stack([entry.update.values for entry in carry.entries], dim=1) + total = gates.sum(1) + if current_gate is not None: + total = total + current_gate + weights = (total.unsqueeze(1) - gates.cumsum(1)).exp() + if gates.ndim == 3: + weighted = keys * weights + weighted = forward_value(weighted, weighted.to(torch.bfloat16)) + return state * total.exp().unsqueeze(-1) + weighted.transpose(-1, -2) @ updates + weighted = updates * weights.unsqueeze(-1) + weighted = forward_value(weighted, weighted.to(torch.bfloat16)) + return state * total.exp()[:, None, None] + keys.transpose(-1, -2) @ weighted + + +def _encode( + value, + enabled, + block_v, + *, + state_format="fp8_e4m3", + state_quantizer=None, +): + """Apply TensorQuantizer state QDQ with straight-through gradients.""" + if not enabled: + return EncodedLinearAttentionTensor(value, None, "identity", None) + if state_quantizer is not None and state_quantizer.block_sizes is not None: + # TensorQuantizer owns dynamic scales; the floating carry needs only its QDQ output. + return EncodedLinearAttentionTensor( + state_quantizer(value), None, state_format, state_quantizer.block_sizes[-1] + ) + decoded, scales = _tile_qdq(value, block_v, state_format, state_quantizer=state_quantizer) + return EncodedLinearAttentionTensor(decoded, scales, state_format, block_v) + + +def _signature(config, state_qdq, block_v, state_format, state_quantizer): + """Record the execution and quantizer settings required when resuming a state.""" + signature = config.model_dump_json() + f"/{state_qdq}/{block_v}/{state_format}" + if state_quantizer is not None and state_quantizer.block_sizes is not None: + signature += f"/group={state_quantizer.block_sizes[-1]}" + return signature + + +def _sum_keys(value): + keys = value.shape[-2] + padded = 1 << (keys - 1).bit_length() + if keys != padded: + value = torch.nn.functional.pad(value, (0, 0, 0, padded - keys)) + while value.shape[-2] > 1: + half = value.shape[-2] // 2 + value = value[..., :half, :] + value[..., half:, :] + return value[..., 0, :] + + +def _prepare_carry( + q, + k, + v, + g, + beta, + config, + state_qdq, + block_v, + initial_state, + carry, + position, + state_format, + state_quantizer, +): + """Initialize from a dense state or validate and resume an existing runtime state. + + The first nonempty call applies the handoff transform/QDQ once. A resumed + state preserves its anchor, replay cursor, and token position; an empty call + leaves initial encoding deferred until a token is actually processed. + """ + if q.ndim != 3 or k.shape != q.shape or v.shape[:2] != q.shape[:2]: + raise ValueError("q/k/v must have aligned [T,H,D] shapes") + if beta.shape != q.shape[:2] or g.shape not in (beta.shape, q.shape): + raise ValueError("beta must be [T,H]; g must be [T,H] or [T,H,Dk]") + if block_v not in (16, 32, 64, 128): + raise ValueError("block_v must be 16, 32, 64, or 128") + if state_format not in ("fp8_e4m3", "int8"): + raise ValueError("State format must be fp8_e4m3 or int8") + hadamard = config.state_codec == "int8_hadamard32" + if hadamard: + if state_quantizer is not None and state_quantizer.block_sizes is not None: + raise ValueError("TensorQuantizer block_sizes requires state_codec='tile'") + if state_qdq and state_format != "int8": + raise ValueError("int8_hadamard32 requires INT8 state quantization") + if v.shape[-1] % 32 or block_v < 32: + raise ValueError("int8_hadamard32 requires Dv divisible by 32 and block_v >= 32") + signature = _signature(config, state_qdq, block_v, state_format, state_quantizer) + if carry is not None and initial_state is not None: + raise ValueError("Supply either carry or initial_state") + shape = (q.shape[1], q.shape[2], v.shape[2]) + if carry is None: + initial_state = q.new_zeros(shape) if initial_state is None else initial_state + carry = LinearAttentionState( + _encode(initial_state, False, block_v, state_format=state_format), + (), + position, + False, + signature, + ) + if carry.signature != signature or carry.anchor.values.shape != shape: + raise ValueError("Carry policy or state shape does not match this recurrence") + if carry.cursor >= config.replay_window: + raise ValueError("Replay cursor must be below the refresh window") + if len(q) and not carry.started: + initial = _hadamard32(carry.anchor.values) if hadamard else carry.anchor.values + if config.precision == "replayssm": + # The optional serving fork owns checkpoint rounding and Hadamard arithmetic. + from ...kernels.quantization.linear_attention.serving.replay import checkpoint + + with torch.no_grad(): + decoded, metadata = checkpoint(carry.anchor.values, state_qdq) + anchor = EncodedLinearAttentionTensor( + forward_value(initial, decoded), + metadata, + "int8" if state_qdq else "identity", + 32 if state_qdq else None, + ) + else: + anchor = _encode( + initial, + state_qdq, + block_v, + state_format=state_format, + state_quantizer=state_quantizer, + ) + carry = LinearAttentionState( + anchor, + (), + carry.position, + True, + signature, + "hadamard32" if hadamard else "identity", + config.precision, + ) + return carry, signature + + +def recurrent_decode( + q, + k, + v, + g, + beta, + *, + config: LinearAttentionConfig, + state_qdq=False, + state_format="fp8_e4m3", + state_quantizer=None, + initial_state=None, + carry: LinearAttentionState | None = None, + position=0, + scale=None, + use_qk_l2norm_in_kernel=False, + replay_gate_inputs=None, +) -> tuple[torch.Tensor, LinearAttentionState]: + """Run a recurrent suffix and return its outputs and resumable runtime state. + + Q/K are [T,H,K], V is [T,H,V], and beta is [T,H]; heads must be aligned. + Log gates are [T,H] for GDN or [T,H,K] for KDA. The first token consumes the + handoff state after enabled QDQ. Each token reads its output from the working + state before checkpoint rounding; replay_window determines when that state + is saved. Inputs and the initial state are promoted to FP32 working values; + native forward values use a differentiable adjoint for training. + + Args: + config: The same LinearAttentionConfig used by the chunked prefix. + initial_state: Dense [H,K,V] state in the original value basis, normally + from prefill. Defaults to zeros; mutually exclusive with ``carry``. + carry: LinearAttentionState returned by an earlier call. Pass it directly + to preserve the replay window and autograd history across calls. + position: Starting token offset when creating a state, usually the prefix + length. An existing carry retains its own position. + + Returns: + Outputs [T,H,V] and the updated LinearAttentionState, both retaining their + graphs. An empty call performs no state write or initial-state QDQ. + """ + state_quantizer, state_qdq, state_format = _resolve_state_quantizer( + state_quantizer, state_qdq, state_format + ) + if config.backend != "serving": + raise ValueError("State QAT requires backend='serving'") + block_v = config.state_block_v + serving = config.precision in ("vllm", "vllm_0_15") + native_replay = config.precision == "replayssm" + if q.device.type != "cuda": + raise ValueError("Serving arithmetic requires CUDA") + # BF16 inputs need FP32 working values for both replay reconstruction and its adjoint. + q, k, v, g, beta = (x.float() for x in (q, k, v, g, beta)) + if initial_state is not None: + initial_state = initial_state.float() + if ( + native_replay + and len(q) + and ( + q.shape[-1] < 32 + or q.shape[-1] & (q.shape[-1] - 1) + or (g.ndim == 2 and replay_gate_inputs is None) + ) + ): + raise ValueError("ReplaySSM requires power-of-two K >= 32 and raw GDN gate inputs") + carry, signature = _prepare_carry( + q, + k, + v, + g, + beta, + config, + state_qdq, + block_v, + initial_state, + carry, + position, + state_format, + state_quantizer, + ) + if len(q) == 0: + # Keep empty input gradients defined without introducing a state write. + zero = (q.sum() + k.sum() + v.sum() + g.sum() + beta.sum()) * 0 + output = v + zero + return output, carry + original_v = v + if config.state_codec == "int8_hadamard32": + v = _hadamard32(v) + scale = q.shape[-1] ** -0.5 if scale is None else scale + outputs = [] + native_outputs = [] + with torch.autocast(device_type=q.device.type, enabled=False): + for t in range(len(q)): + entries = carry.entries + state = carry.reconstruct(original_basis=False) + gate = g[t] + if native_replay: + from ...kernels.quantization.linear_attention.serving.replay import step + + with torch.no_grad(): + native = step( + q[t], + k[t], + original_v[t], + gate, + beta[t], + carry, + config.replay_window, + state_qdq, + scale, + use_qk_l2norm_in_kernel, + ( + replay_gate_inputs[0][t], + replay_gate_inputs[1][t], + *replay_gate_inputs[2:], + ) + if replay_gate_inputs is not None + else None, + ) + native_outputs.append(native[0]) + if serving: + from ._vllm_autograd import step as serving_step + + native_output, working = serving_step( + q[t], k[t], v[t], gate, beta[t], state, scale, use_qk_l2norm_in_kernel + ) + else: + decay = gate.exp().unsqueeze(-1) + if gate.ndim == 1: + decay = decay.unsqueeze(-1) + current_key = k[t] + if native_replay and use_qk_l2norm_in_kernel: + current_key = ( + current_key / (current_key.square().sum(-1, keepdim=True) + 1e-6).sqrt() + ) + key = EncodedLinearAttentionTensor(current_key, None, "identity", None) + decayed = state * decay + if native_replay and gate.ndim == 2: + decayed = _replay_state(carry, gate) + residual = v[t] - _sum_keys(key.values.unsqueeze(-1) * decayed) + update = EncodedLinearAttentionTensor( + beta[t].unsqueeze(-1) * residual, None, "identity", None + ) + working = decayed + key.values.unsqueeze(-1) * update.values.unsqueeze(-2) + if config.replay_window > 1: + if native_replay and carry.cursor + 1 < config.replay_window: + key.values = forward_value(key.values, native[3]) + update.values = forward_value(update.values, native[4]) + gate = forward_value(gate, native[5]) + entries = (*entries, ReplayEntry(key, update, gate)) + # Window 1 checkpoints every token; replay checkpoints only at a full window. + # Between checkpoints, keep the anchor and save the native BF16 update values. + refresh = config.replay_window == 1 or len(entries) == config.replay_window + if refresh: + if native_replay: + # Native step already quantized the checkpoint; attach its values to the adjoint. + anchor = EncodedLinearAttentionTensor( + forward_value(working, native[1]), + native[2], + "int8" if state_qdq else "identity", + 32 if state_qdq else None, + ) + else: + anchor = _encode( + working, + state_qdq, + block_v, + state_format=state_format, + state_quantizer=state_quantizer, + ) + entries = () + else: + anchor = carry.anchor + # A new container keeps tensors needed by earlier tokens' backward passes intact. + next_carry = LinearAttentionState( + anchor, + entries, + carry.position + 1, + True, + signature, + carry.value_basis, + config.precision, + ) + query = q[t] + if native_replay and use_qk_l2norm_in_kernel: + query = query / (query.square().sum(-1, keepdim=True) + 1e-6).sqrt() + # Read this token from the working state, before the next-state checkpoint QDQ. + outputs.append( + native_output if serving else _sum_keys(query.unsqueeze(-1) * working) * scale + ) + carry = next_carry + output = torch.stack(outputs) + if config.state_codec == "int8_hadamard32": + output = _hadamard32(output) + if native_replay: + output = forward_value(output, torch.stack(native_outputs)) + return output, carry diff --git a/modelopt/torch/quantization/linear_attention/gdn.py b/modelopt/torch/quantization/linear_attention/gdn.py new file mode 100644 index 00000000000..7b73724dec7 --- /dev/null +++ b/modelopt/torch/quantization/linear_attention/gdn.py @@ -0,0 +1,96 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""GDN adapter for serving-aligned recurrent-state QAT.""" + +from .config import LinearAttentionConfig +from .training import _prefill_decode_forward, _prepare_prefill_inputs + +__all__ = ["gdn_state_qat", "matmul_gdn"] + + +def gdn_state_qat( + q, + k, + v, + g, + beta, + *, + policy: LinearAttentionConfig, + state_qdq=False, + state_format="fp8_e4m3", + state_quantizer=None, + scale=None, + initial_state=None, + output_final_state=False, + use_qk_l2norm_in_kernel=False, + use_gate_in_kernel=False, + use_beta_sigmoid_in_kernel=False, + allow_neg_eigval=False, + A_log=None, # noqa: N803 - match the FLA kernel signature + dt_bias=None, + cu_seqlens=None, + cu_seqlens_cpu=None, + state_v_first=False, + chunk_size=64, + cp_context=None, + prefill_lengths=None, + replay_gate_inputs=None, +): + """Adapt Megatron's FLA-style GDN call to serving-aligned state QAT. + + Validate prepared scalar log gates and promote inputs to FP32 working values. + The shared training forward runs the chunked prefix and recurrent suffix with + native forward values, configured state QDQ, and a differentiable Torch adjoint. + """ + if cp_context is not None: + raise NotImplementedError("GDN state QAT does not support context parallelism") + output_dtype = q.dtype + beta_dtype = beta.dtype + q, k, v, g, beta = _prepare_prefill_inputs( + q, k, v, g, beta, policy=policy, chunk_size=chunk_size + ) + if use_gate_in_kernel or use_beta_sigmoid_in_kernel: + raise ValueError( + "Serving GDN expects prepared log gates and beta from the Megatron adapter" + ) + if g.ndim != 3: + raise ValueError("gdn_state_qat requires scalar GDN log gates") + return _prefill_decode_forward( + q, + k, + v, + g, + beta, + policy=policy, + state_qdq=state_qdq, + state_format=state_format, + state_quantizer=state_quantizer, + scale=scale, + initial_state=initial_state, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + cu_seqlens_cpu=cu_seqlens_cpu, + state_v_first=state_v_first, + output_dtype=output_dtype, + beta_dtype=beta_dtype, + prefill_lengths=prefill_lengths, + replay_gate_inputs=replay_gate_inputs, + use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, + ) + + +# Compatibility name for callers using the original adapter API. +matmul_gdn = gdn_state_qat diff --git a/modelopt/torch/quantization/linear_attention/kda.py b/modelopt/torch/quantization/linear_attention/kda.py new file mode 100644 index 00000000000..f2238c2c7ad --- /dev/null +++ b/modelopt/torch/quantization/linear_attention/kda.py @@ -0,0 +1,122 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""KDA adapter for serving-aligned recurrent-state QAT.""" + +import torch +import torch.nn.functional as F + +from .training import _prefill_decode_forward, _prepare_prefill_inputs +from .utils import forward_value + +__all__ = ["kda_state_qat", "matmul_kda"] + + +def kda_state_qat( + q, + k, + v, + g, + beta, + *, + policy, + state_qdq=False, + state_format="fp8_e4m3", + state_quantizer=None, + scale=None, + initial_state=None, + output_final_state=False, + use_qk_l2norm_in_kernel=False, + use_gate_in_kernel=False, + use_beta_sigmoid_in_kernel=False, + allow_neg_eigval=False, + A_log=None, # noqa: N803 - match the FLA kernel signature + dt_bias=None, + safe_gate=False, + lower_bound=None, + cu_seqlens=None, + cu_seqlens_cpu=None, + state_v_first=False, + chunk_size=64, + cp_context=None, + disable_recompute=False, + return_intermediate_states=False, + prefill_lengths=None, +): + """Adapt Megatron's FLA-style KDA call to serving-aligned state QAT. + + Prepare per-key-channel gates and optional beta activation using FLA's formula. + The shared training forward runs the chunked prefix and recurrent suffix with + native forward values, configured state QDQ, and a differentiable Torch adjoint. + """ + if cp_context is not None or disable_recompute or return_intermediate_states: + raise NotImplementedError( + "KDA state QAT does not support CP or FLA recompute/intermediate flags" + ) + if allow_neg_eigval and not use_beta_sigmoid_in_kernel: + raise ValueError("allow_neg_eigval requires use_beta_sigmoid_in_kernel") + if lower_bound is not None or safe_gate: + raise ValueError("Serving arithmetic uses the native softplus KDA gate") + output_dtype = q.dtype + beta_dtype = beta.dtype + q, k, v, g, beta = _prepare_prefill_inputs( + q, k, v, g, beta, policy=policy, chunk_size=chunk_size + ) + dtype = q.dtype + if g.ndim != 4: + raise ValueError("kda_state_qat requires per-key-channel KDA log gates") + if use_gate_in_kernel: + raw_gate = g + if A_log is None: + raise ValueError("Fused KDA gate requires A_log") + if dt_bias is not None: + g = g + dt_bias.to(dtype).reshape(g.shape[-2:]) + rate = A_log.to(dtype).exp().reshape(g.shape[-2], 1) + g = -rate * F.softplus(g) + # Import the optional vLLM backend only for the native precision profile. + from ...kernels.quantization.linear_attention.serving.forward import fused_kda_gate + + with torch.no_grad(): + native_gate = fused_kda_gate( + raw_gate.flatten(-2).contiguous(), A_log, raw_gate.shape[-1], g_bias=dt_bias + ) + g = forward_value(g, native_gate) + if use_beta_sigmoid_in_kernel: + beta = beta.sigmoid() * (2.0 if allow_neg_eigval else 1.0) + return _prefill_decode_forward( + q, + k, + v, + g, + beta, + policy=policy, + state_qdq=state_qdq, + state_format=state_format, + state_quantizer=state_quantizer, + scale=scale, + initial_state=initial_state, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + cu_seqlens_cpu=cu_seqlens_cpu, + state_v_first=state_v_first, + output_dtype=output_dtype, + beta_dtype=beta_dtype, + prefill_lengths=prefill_lengths, + use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, + ) + + +# Compatibility name for callers using the original adapter API. +matmul_kda = kda_state_qat diff --git a/modelopt/torch/quantization/linear_attention/training.py b/modelopt/torch/quantization/linear_attention/training.py new file mode 100644 index 00000000000..ed8401d5107 --- /dev/null +++ b/modelopt/torch/quantization/linear_attention/training.py @@ -0,0 +1,227 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Training forwards with a chunked prefill prefix and recurrent decode suffix.""" + +from contextlib import contextmanager +from itertools import pairwise + +import torch + +import modelopt.torch.utils.distributed as dist + +from .decode import recurrent_decode +from .utils import _resolve_state_quantizer, _state_qdq, forward_value + +__all__ = ["linear_attention_training_phase"] + + +def _lengths(values): + if isinstance(values, torch.Tensor): + values = values.tolist() + values = tuple(values) + if any(type(value) is not int or value < 0 for value in values): + raise ValueError("Prefill lengths must be nonnegative integers") + return values + + +@contextmanager +def linear_attention_training_phase(model, prefill_lengths): + """Supply explicit sequence phases through forward and checkpointed backward. + + Runtime phase metadata is local to the converted layers and restored on exit. + Keep this context active through backward when activation checkpointing is used. + """ + # The plugin imports this numerical package during quantization initialization. + from ..plugins.linear_attention import _LinearAttentionQuantMixin + + lengths = _lengths(prefill_lengths) + layers = [ + m + for m in model.modules() + if isinstance(m, _LinearAttentionQuantMixin) + and m.linear_attention_config.backend == "serving" + ] + # A pipeline stage may have no selected layers; conversion validates global matches. + if not layers and dist.size() == 1: + raise ValueError("The model has no converted decode-aware linear-attention layers") + previous = [getattr(m, "_linear_attention_prefill_lengths", None) for m in layers] + try: + for module in layers: + module._linear_attention_prefill_lengths = lengths + yield model + finally: + for module, original in zip(layers, previous): + module._linear_attention_prefill_lengths = original + + +def _prepare_prefill_inputs(q, k, v, g, beta, *, policy, chunk_size): + """Validate the native policy and prepare FP32 values for the training adjoint.""" + if policy.backend != "serving" or chunk_size != 64: + raise ValueError("State training requires backend='serving' and chunk_size=64") + if q.device.type != "cuda" or any(x.dtype != torch.bfloat16 for x in (q, k, v)): + raise ValueError("Serving arithmetic requires CUDA BF16 Q/K/V inputs") + return tuple(x.float() for x in (q, k, v, g, beta)) + + +def _prefill_decode_forward( + q, + k, + v, + g, + beta, + *, + policy, + state_qdq, + state_format, + state_quantizer, + scale, + initial_state, + output_final_state, + cu_seqlens, + cu_seqlens_cpu, + state_v_first, + output_dtype, + beta_dtype, + prefill_lengths, + use_qk_l2norm_in_kernel=False, + replay_gate_inputs=None, +): + """Run both prefill and decode phases in one differentiable training forward. + + Each sequence's chunked prefix produces the state for its token or ReplaySSM + suffix. ``recurrent_decode`` wraps that dense state in LinearAttentionState + and preserves its gradient connection to prefill. Outputs are joined in token + order; the final runtime state is reconstructed to the caller's dense layout. + + Args: + prefill_lengths: Prefix token count per sequence. For 128 tokens, a value + of 64 selects 64 chunked prefill tokens followed by 64 recurrent tokens. + """ + state_quantizer, state_qdq, state_format = _resolve_state_quantizer( + state_quantizer, state_qdq, state_format + ) + if prefill_lengths is None: + raise ValueError("Decode-aware training requires explicit per-sequence prefill lengths") + if q.ndim != 4 or k.shape != q.shape or v.ndim != 4 or v.shape[:2] != q.shape[:2]: + raise ValueError("q/k and v must have compatible [B,T,H,D] shapes") + batch, length, key_heads, keys = q.shape + heads, values = v.shape[2:] + if batch < 1 or key_heads < 1 or heads % key_heads or beta.shape != (batch, length, heads): + raise ValueError("Invalid batch/head dimensions or beta shape") + if g.shape not in (beta.shape, (*beta.shape, keys)): + raise ValueError("Invalid GDN/KDA log-retention shape") + q, k = (x.repeat_interleave(heads // key_heads, dim=2) for x in (q, k)) + boundaries = cu_seqlens_cpu if cu_seqlens_cpu is not None else cu_seqlens + if boundaries is None: + sequences = [(b, 0, length) for b in range(batch)] + else: + bounds = _lengths(boundaries) + if ( + batch != 1 + or len(bounds) < 2 + or bounds[0] != 0 + or bounds[-1] != length + or any(a > b for a, b in pairwise(bounds)) + ): + raise ValueError( + "Packed boundaries must partition a batch of one, allowing empty entries" + ) + sequences = [(0, a, b) for a, b in pairwise(bounds)] + prefixes = _lengths(prefill_lengths) + if len(prefixes) != len(sequences) or any( + p > end - start for p, (_, start, end) in zip(prefixes, sequences) + ): + raise ValueError("Supply one valid prefill length per sequence") + if ( + keys > 256 + or (g.ndim == 4 and values != keys) + or (use_qk_l2norm_in_kernel and keys & (keys - 1)) + ): + raise ValueError( + "Serving arithmetic requires K <= 256, KDA V=K, and power-of-two K for normalization" + ) + if initial_state is None: + states = q.new_zeros(len(sequences), heads, keys, values) + else: + states = initial_state.transpose(-1, -2) if state_v_first else initial_state + states = states.to(q.dtype) + if states.shape != (len(sequences), heads, keys, values): + raise ValueError("Initial state shape does not match sequence/head dimensions") + outputs, finals = [], [] + for n, (b, start, end) in enumerate(sequences): + split = start + prefixes[n] + prefix, state = q.new_empty(0, heads, values), states[n] + if prefixes[n]: + with torch.autocast(device_type=q.device.type, enabled=False): + from ._vllm_autograd import prefix as serving_prefix + + # A continuation prefill consumes a stored cache just as native serving + # does. A fresh zero-state prefix has no incoming cache to quantize. + if initial_state is not None and state_qdq: + if policy.precision == "replayssm": + from ...kernels.quantization.linear_attention.serving.replay import ( + checkpoint, + original_basis, + ) + + with torch.no_grad(): + decoded, _ = checkpoint(state, True) + decoded = original_basis(decoded) + state = forward_value(state, decoded) + else: + state = _state_qdq( + state, policy.state_block_v, state_format, state_quantizer + ) + prefix, state = serving_prefix( + q=q[b, start:split], + k=k[b, start:split], + v=v[b, start:split], + g=g[b, start:split], + beta=beta[b, start:split], + state=state, + scale=keys**-0.5 if scale is None else scale, + beta_dtype=beta_dtype, + normalize=use_qk_l2norm_in_kernel, + ) + # Keep the prefix state attached so suffix losses backpropagate through prefill. + # recurrent_decode applies configured state QDQ at the handoff and suffix writes. + suffix, carry = recurrent_decode( + *(x[b, split:end] for x in (q, k, v, g, beta)), + config=policy, + state_qdq=state_qdq, + state_format=state_format, + state_quantizer=state_quantizer, + initial_state=state, + position=prefixes[n], + replay_gate_inputs=( + ( + replay_gate_inputs[0][b, split:end], + replay_gate_inputs[1][b, split:end], + *replay_gate_inputs[2:], + ) + if replay_gate_inputs is not None + else None + ), + use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, + scale=scale, + ) + outputs.append(torch.cat((prefix, suffix))) + finals.append(carry.reconstruct()) + output = torch.stack(outputs) if boundaries is None else torch.cat(outputs).unsqueeze(0) + final = torch.stack(finals) + if state_v_first: + final = final.transpose(-1, -2) + return output.to(output_dtype), final if output_final_state else None diff --git a/modelopt/torch/quantization/linear_attention/utils.py b/modelopt/torch/quantization/linear_attention/utils.py index 4f1e18a174d..afb40b73c5c 100644 --- a/modelopt/torch/quantization/linear_attention/utils.py +++ b/modelopt/torch/quantization/linear_attention/utils.py @@ -15,19 +15,59 @@ """Shared helpers for linear-attention quantization.""" -from ..nn import TensorQuantizer +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +if TYPE_CHECKING: + from ..nn import TensorQuantizer __all__ = [] +_STATE_FORMATS: dict[int | tuple[int, int], str] = {(4, 3): "fp8_e4m3", 8: "int8"} + + +class _ForwardValue(torch.autograd.Function): + @staticmethod + def forward(ctx, value, rounded): + return rounded.to(value.dtype) + + @staticmethod + def backward(ctx, grad): + return grad, None + + +def forward_value(value, rounded): + # Unlike value + (rounded - value).detach(), this cannot lose low bits by cancellation. + return _ForwardValue.apply(value, rounded) + + +def validate_gdn_quantizer( + quantizer: TensorQuantizer, + *, + name: str, + num_bits: tuple[int | tuple[int, int], ...] = ((4, 3),), + block_sizes: tuple[int, ...] = (), +) -> None: + """Check supported formats and the custom backward's identity STE.""" + # Numerical helpers load while QuantizeConfig initializes; defer the quantizer import. + from ..nn import TensorQuantizer -def validate_gdn_quantizer(quantizer: TensorQuantizer, *, name: str) -> None: - """Check the dynamic E4M3 contract and the custom backward's identity STE.""" if not isinstance(quantizer, TensorQuantizer): raise ValueError(f"{name} requires a single TensorQuantizer") if not ( quantizer._dynamic - and quantizer.num_bits == (4, 3) - and quantizer.block_sizes is None + and quantizer.num_bits in num_bits + and (quantizer.num_bits != 8 or (not quantizer.unsigned and quantizer.narrow_range)) + and ( + quantizer.block_sizes is None + or ( + quantizer.num_bits == 8 + and quantizer.block_sizes in ({-1: size} for size in block_sizes) + ) + ) and quantizer.fake_quant and quantizer._pass_through_bwd and not quantizer.rotate_is_enabled @@ -37,7 +77,94 @@ def validate_gdn_quantizer(quantizer: TensorQuantizer, *, name: str) -> None: and not quantizer._use_constant_amax ): raise ValueError( - f"{name} supports only dynamic E4M3 fake quantization with " - "pass_through_bwd=True, no block_sizes, rotation, pre-scaling, bias, constant " + f"{name} supports only dynamic fake quantization with num_bits in {num_bits}, " + f"INT8 block_sizes in {block_sizes} (or no blocks), " + "pass_through_bwd=True, no rotation, pre-scaling, bias, constant " "amax, or custom backend. Other gradient rules and formats are not implemented." ) + + +def state_quantizer_config( + quantizer: TensorQuantizer, *, name="state_quantizer" +) -> tuple[str, int]: + """Validate the state quantizer and derive its format and last-axis group size. + + A zero group size retains the legacy policy's full-key state tiles. + """ + validate_gdn_quantizer( + quantizer, name=name, num_bits=tuple(_STATE_FORMATS), block_sizes=(16, 32, 64) + ) + if quantizer.block_sizes is not None: + return _STATE_FORMATS[quantizer.num_bits], quantizer.block_sizes[-1] + if quantizer.axis != (0, 1): + raise ValueError(f"{name} supports only axis=(0, 1) with state_block_v tiling") + return _STATE_FORMATS[quantizer.num_bits], 0 + + +def _make_state_quantizer(state_format): + """Build a standalone-call quantizer; converted modules supply their registered instance.""" + # QuantizeConfig imports this module before the quantizer classes are initialized. + from ..config import QuantizerAttributeConfig + from ..nn import TensorQuantizer + + if state_format not in _STATE_FORMATS.values(): + raise ValueError("State format must be fp8_e4m3 or int8") + return TensorQuantizer( + QuantizerAttributeConfig( + num_bits=(4, 3) if state_format == "fp8_e4m3" else 8, + type="dynamic", + axis=(0, 1), + narrow_range=True, + pass_through_bwd=True, + ) + ) + + +def _resolve_state_quantizer(state_quantizer, state_qdq, state_format): + """Resolve legacy flags and registered quantizer settings at either training entry point.""" + if state_quantizer is None and state_qdq: + state_quantizer = _make_state_quantizer(state_format) + if state_quantizer is not None: + state_qdq = state_quantizer.is_enabled and state_quantizer._if_quant + if state_quantizer.is_enabled: + state_format, _ = state_quantizer_config(state_quantizer) + return state_quantizer, state_qdq, state_format + + +def _state_qdq( + state: torch.Tensor, + block_v: int = 64, + state_format: str = "fp8_e4m3", + state_quantizer: TensorQuantizer | None = None, +): + """Dynamic state-tile QDQ with detached scales and identity STE.""" + if state_quantizer is not None and state_quantizer.block_sizes is not None: + return state_quantizer(state) + if state_format not in ("fp8_e4m3", "int8"): + raise ValueError("State format must be fp8_e4m3 or int8") + if block_v not in (16, 32, 64, 128): + raise ValueError("block_v must be 16, 32, 64, or 128") + quantized, _ = _tile_qdq(state, block_v, state_format, state_quantizer=state_quantizer) + return quantized + + +def _tile_qdq(value, block_v, state_format, *, state_quantizer=None): + """Return TensorQuantizer tile QDQ with identity STE and detached scales.""" + quantizer = ( + state_quantizer if state_quantizer is not None else _make_state_quantizer(state_format) + ) + rounded, scales = [], [] + for part in value.split(block_v, dim=-1): + # Canonical [N, H, K*BV] shape preserves per-head tile scales for both state ranks. + tensor = part.flatten(-2) + inputs = tensor.reshape(1, -1, tensor.shape[-1]) + decoded = quantizer(inputs) + amax = quantizer._get_amax(inputs).float() + if quantizer.num_bits == 8: + scale = amax / quantizer.maxbound + else: + safe_amax = torch.where(amax <= 2**-24, torch.ones_like(amax), amax) + scale = torch.div(quantizer.maxbound, safe_amax).reciprocal() + rounded.append(decoded.reshape_as(part)) + scales.append(scale.reshape(tensor.shape[:-1])) + return torch.cat(rounded, dim=-1).to(value.dtype), torch.stack(scales, dim=-1) diff --git a/modelopt/torch/quantization/model_quant.py b/modelopt/torch/quantization/model_quant.py index 5b53526aba6..3266182bacf 100644 --- a/modelopt/torch/quantization/model_quant.py +++ b/modelopt/torch/quantization/model_quant.py @@ -33,7 +33,6 @@ from modelopt.torch.opt.utils import forward_with_reshard from modelopt.torch.quantization.config import QuantizeConfig from modelopt.torch.quantization.conversion import ( - _validate_linear_attention_quantizers, preserve_quantizer_attributes_context, set_quantizer_attributes_partial, set_quantizer_by_cfg, @@ -390,9 +389,16 @@ def forward_loop(model) -> None: if not is_quantized(model): model = apply_mode(model, mode=[("quantize", dict(config))], registry=QuantizeModeRegistry) else: + # Defer shared plugin imports to avoid the plugin/conversion initialization cycle. + from .plugins.linear_attention import ( + _apply_linear_attention_policy, + _validate_linear_attention, + ) + # Already quantized, so lets apply the quant_cfg from the config set_quantizer_by_cfg(model, quantize_config.quant_cfg) - _validate_linear_attention_quantizers(model) + _apply_linear_attention_policy(model, quantize_config) + _validate_linear_attention(model) # Fail before calibration rather than after exporting an unquantized checkpoint. _check_weight_quantization_took_effect(model, quantize_config) _check_indexer_quantization_took_effect(model, quantize_config) diff --git a/modelopt/torch/quantization/nn/modules/tensor_quantizer.py b/modelopt/torch/quantization/nn/modules/tensor_quantizer.py index a698514d53a..5e3b677ea2f 100644 --- a/modelopt/torch/quantization/nn/modules/tensor_quantizer.py +++ b/modelopt/torch/quantization/nn/modules/tensor_quantizer.py @@ -999,9 +999,12 @@ def set_quant_params(axis, block_reshape_size, padding, slices, amax_shape=None) if amax_shape: self._amax_shape_for_export = amax_shape - # Reshape size have already been set - if hasattr(self, "_block_reshape_size"): + # Static scales require fixed groups; dynamic scales allow changing input shapes. + if hasattr(self, "_block_reshape_size") and not self._dynamic: return + for attribute in ("_padding", "_slices"): + if hasattr(self, attribute): + delattr(self, attribute) reshape_size, quantize_axis, paddings, slices = [], [], [], [] diff --git a/modelopt/torch/quantization/plugins/__init__.py b/modelopt/torch/quantization/plugins/__init__.py index 7234ca81f76..247dacdd909 100644 --- a/modelopt/torch/quantization/plugins/__init__.py +++ b/modelopt/torch/quantization/plugins/__init__.py @@ -39,7 +39,7 @@ from .attention import * from .custom import * -from .gated_delta_net import * +from .gdn import * with import_plugin("diffusers"): from .diffusion.diffusers import * @@ -72,6 +72,7 @@ with import_plugin("vllm"): from .vllm import * from .vllm_indexer import * + from .vllm_linear_attention import * with import_plugin("trl"): from .trl import * diff --git a/modelopt/torch/quantization/plugins/gated_delta_net.py b/modelopt/torch/quantization/plugins/gated_delta_net.py deleted file mode 100644 index df70683cc2f..00000000000 --- a/modelopt/torch/quantization/plugins/gated_delta_net.py +++ /dev/null @@ -1,128 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Fake quantization of the GatedDeltaNet (GDN) recurrent state. - -The chunked gated-delta-rule kernel keeps each head's ``[K, V]`` recurrent state in fp32 inside -one Triton launch and carries it from chunk to chunk. To emulate a deployment that stores that -state in FP8, ModelOpt runs an adapted copy of the kernel -(:mod:`modelopt.torch.kernels.quantization.linear_attention`) that fake-quantizes the state to -E4M3 at the end of every chunk, with a scale computed inside the kernel from the state itself. -The backward pass recomputes the same quantized states and passes the state gradient straight -through the quantization, so QAT and QAD train against the quantized recurrence. A second -quantizer covers ``w``, the WY-transformed keys that multiply the state; ``w`` is a regular tensor, -so it is fake-quantized by the ``TensorQuantizer`` itself before the kernel reads it. -""" - -from collections.abc import Callable -from functools import partial -from typing import Any - -import torch - -from ..config import QuantizerAttributeConfig -from ..linear_attention.utils import validate_gdn_quantizer -from ..nn import QuantModule, TensorQuantizer - -__all__ = ["GatedDeltaNetStateQuantMixin"] - -GatedDeltaRuleFn = Callable[..., tuple[torch.Tensor, torch.Tensor | None]] - - -def _fla_chunk_gated_delta_rule() -> GatedDeltaRuleFn: - # FLA is an optional, heavy dependency needed only when quantization is enabled. - from fla.ops.gated_delta_rule import chunk_gated_delta_rule - - return chunk_gated_delta_rule - - -def _state_qdq_chunk_gated_delta_rule() -> GatedDeltaRuleFn: - # Imported on first use: flash-linear-attention is a heavy optional dependency that only the - # enabled quantizer needs, and importing it warns on machines without a GPU. - try: - from modelopt.torch.kernels.quantization.linear_attention.fla_chunk_gated_delta_rule import ( - chunk_gated_delta_rule, - ) - except ImportError as e: - raise RuntimeError( - "GDN fake quantization needs Triton and fla-core==0.5.1 on a CUDA " - f"device; importing the state-quantizing kernel failed with {e!r}." - ) from e - return chunk_gated_delta_rule - - -class GatedDeltaNetStateQuantMixin(QuantModule): - """Adds ``gdn_state_quantizer`` and ``gdn_w_quantizer`` to a GatedDeltaNet module. - - Subclasses route the module's chunked gated-delta-rule call through - :meth:`_state_quantized_chunk_gated_delta_rule`. Both quantizers start disabled; enable them - with ``quant_cfg`` entries on ``*gdn_state_quantizer`` / ``*gdn_w_quantizer`` such as the - ``configs/ptq/units/gdn_state_fp8_dynamic`` and ``gdn_w_fp8_dynamic`` recipe units. The state - quantizer carries the fused QDQ configuration. Both sites currently require dynamic - E4M3 and identity STE. State QDQ uses fixed 64-column value tiles and 64-token chunks. - """ - - def _setup(self): - for name in ("gdn_state_quantizer", "gdn_w_quantizer"): - self._register_temp_attribute( - name, TensorQuantizer(QuantizerAttributeConfig(enable=False)) - ) - - def validate_linear_attention(self) -> None: - """Reject numerical settings that the fused training path cannot implement.""" - for name in ("gdn_state_quantizer", "gdn_w_quantizer"): - quantizer = getattr(self, name) - if quantizer.is_enabled: - validate_gdn_quantizer(quantizer, name=name) - # The state handle configures fused QDQ; W's grouping is executed by TensorQuantizer. - if self.gdn_state_quantizer.is_enabled and self.gdn_state_quantizer.axis != (0, 1): - raise ValueError( - "gdn_state_quantizer supports only axis=(0, 1) with 64-column value tiling" - ) - - def modelopt_post_restore(self, prefix: str = ""): - """Validate restored quantizers against the fused training kernel requirements.""" - super().modelopt_post_restore(prefix) - self.validate_linear_attention() - - def _state_quantized_chunk_gated_delta_rule( - self, gated_delta_rule: GatedDeltaRuleFn, *args: Any, **kwargs: Any - ) -> tuple[torch.Tensor, torch.Tensor | None]: - """Call ``gated_delta_rule`` or, if a quantizer is on, the vendored quantizing copy.""" - self.validate_linear_attention() - quantize_state = self.gdn_state_quantizer.is_enabled and self.gdn_state_quantizer._if_quant - quantize_w = self.gdn_w_quantizer.is_enabled - if not (quantize_state or quantize_w): - return gated_delta_rule(*args, **kwargs) - while isinstance(gated_delta_rule, partial): - args = (*gated_delta_rule.args, *args) - kwargs = {**gated_delta_rule.keywords, **kwargs} - gated_delta_rule = gated_delta_rule.func - if gated_delta_rule is not _fla_chunk_gated_delta_rule(): - raise NotImplementedError( - "GatedDeltaNet quantization supports only FLA's chunk_gated_delta_rule " - f"callable or a functools.partial of it; got {gated_delta_rule!r}." - ) - chunk_size = kwargs.pop("chunk_size", 64) - if chunk_size != 64: - raise ValueError("GDN fake quantization supports only chunk_size=64") - return _state_qdq_chunk_gated_delta_rule()( - *args, - chunk_size=chunk_size, - state_qdq=int(quantize_state), - state_qdq_block_v=64, - w_quantizer=self.gdn_w_quantizer if quantize_w else None, - **kwargs, - ) diff --git a/modelopt/torch/quantization/plugins/gdn.py b/modelopt/torch/quantization/plugins/gdn.py new file mode 100644 index 00000000000..d1cdc04b7c5 --- /dev/null +++ b/modelopt/torch/quantization/plugins/gdn.py @@ -0,0 +1,95 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""GatedDeltaNet recurrent-state fake quantization. + +State QAT uses an explicit prefix/suffix policy with native serving arithmetic. +""" + +from collections.abc import Callable +from functools import partial +from typing import Any + +import torch + +from ..linear_attention.gdn import gdn_state_qat +from .linear_attention import _LinearAttentionQuantMixin + +__all__ = ["GatedDeltaNetStateQuantMixin"] + +GatedDeltaRuleFn = Callable[..., tuple[torch.Tensor, torch.Tensor | None]] + + +def _fla_chunk_gated_delta_rule() -> GatedDeltaRuleFn: + # FLA is an optional, heavy dependency needed only when quantization is enabled. + from fla.ops.gated_delta_rule import chunk_gated_delta_rule + + return chunk_gated_delta_rule + + +class GatedDeltaNetStateQuantMixin(_LinearAttentionQuantMixin): + """Adds recurrent-state fake quantization to a GatedDeltaNet module. + + Subclasses route the module's chunked gated-delta-rule call through + :meth:`_state_quantized_chunk_gated_delta_rule`. Enable ``*gdn_state_quantizer`` + through ``quant_cfg`` and select an explicit serving policy. State supports dynamic + E4M3 or signed narrow-range INT8 with identity STE. The execution policy is saved in + ModelOpt metadata. ``gdn_w_quantizer`` is a disabled legacy checkpoint handle; + enabling WY operand quantization is no longer supported. + """ + + linear_attention_quantizer_names = ("gdn_state_quantizer", "gdn_w_quantizer") + + def validate_linear_attention(self): + """Reject retired W QAT and validate the state quantizer.""" + if self.gdn_w_quantizer.is_enabled: + raise ValueError( + "GDN W quantization is no longer supported. Disable gdn_w_quantizer; " + "use gdn_state_quantizer with an explicit serving policy for state QAT." + ) + super().validate_linear_attention() + + @property + def gdn_state_qdq_block_v(self) -> int: + """Execution tile width; also sets grouping for legacy tile quantizers.""" + return self.linear_attention_config.state_block_v + + def _state_quantized_chunk_gated_delta_rule( + self, gated_delta_rule: GatedDeltaRuleFn, *args: Any, **kwargs: Any + ) -> tuple[torch.Tensor, torch.Tensor | None]: + """Route enabled state quantization to its training implementation.""" + self.validate_linear_attention() + if not self.linear_attention_is_enabled: + return gated_delta_rule(*args, **kwargs) + while isinstance(gated_delta_rule, partial): + args = (*gated_delta_rule.args, *args) + kwargs = {**gated_delta_rule.keywords, **kwargs} + gated_delta_rule = gated_delta_rule.func + if gated_delta_rule is not _fla_chunk_gated_delta_rule(): + raise NotImplementedError( + "GatedDeltaNet quantization supports only FLA's chunk_gated_delta_rule " + f"callable or a functools.partial of it; got {gated_delta_rule!r}." + ) + chunk_size = kwargs.pop("chunk_size", 64) + if chunk_size != 64: + raise ValueError("GDN fake quantization supports only chunk_size=64") + return gdn_state_qat( + *args, + policy=self.linear_attention_config, + state_quantizer=self.gdn_state_quantizer, + chunk_size=chunk_size, + prefill_lengths=self._linear_attention_prefill_lengths, + **kwargs, + ) diff --git a/modelopt/torch/quantization/plugins/kda.py b/modelopt/torch/quantization/plugins/kda.py new file mode 100644 index 00000000000..04766a5d8e9 --- /dev/null +++ b/modelopt/torch/quantization/plugins/kda.py @@ -0,0 +1,50 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Kernel routing for Megatron Kimi Delta Attention quantization.""" + +from functools import partial + +from ..linear_attention.kda import kda_state_qat +from .linear_attention import _LinearAttentionQuantMixin + +__all__ = ["KimiDeltaAttentionStateQuantMixin"] + + +class KimiDeltaAttentionStateQuantMixin(_LinearAttentionQuantMixin): + """Adds state quantizers and decode-aware kernel routing to Kimi Delta Attention.""" + + linear_attention_quantizer_names = ("kda_state_quantizer",) + + def _state_quantized_chunk_kda(self, kernel, *args, **kwargs): + self.validate_linear_attention() + if not self.linear_attention_is_enabled: + return kernel(*args, **kwargs) + # FLA kernels are an optional dependency; the training layer belongs to Megatron. + from fla.ops.kda import chunk_kda + + while isinstance(kernel, partial): + args = (*kernel.args, *args) + kwargs = {**kernel.keywords, **kwargs} + kernel = kernel.func + if kernel is not chunk_kda: + raise NotImplementedError("KDA quantization requires FLA's chunk_kda callable") + return kda_state_qat( + *args, + policy=self.linear_attention_config, + state_quantizer=self.kda_state_quantizer, + prefill_lengths=self._linear_attention_prefill_lengths, + **kwargs, + ) diff --git a/modelopt/torch/quantization/plugins/linear_attention.py b/modelopt/torch/quantization/plugins/linear_attention.py new file mode 100644 index 00000000000..d5945b57bab --- /dev/null +++ b/modelopt/torch/quantization/plugins/linear_attention.py @@ -0,0 +1,139 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Shared module policy and checkpoint support for linear-attention QAT.""" + +import fnmatch + +import torch + +import modelopt.torch.utils.distributed as dist +from modelopt.torch.opt.conversion import ApplyModeError +from modelopt.torch.utils import get_unwrapped_name + +from ..config import QuantizerAttributeConfig +from ..linear_attention.config import LinearAttentionConfig +from ..linear_attention.utils import state_quantizer_config +from ..nn import QuantModule, TensorQuantizer + +__all__ = [] + + +def _linear_attention_modules(model): + return { + get_unwrapped_name(name, model): module + for name, module in model.named_modules() + if isinstance(module, _LinearAttentionQuantMixin) + } + + +def _apply_linear_attention_policy(model, config): + """Assign complete policies in rule order, before compatibility validation.""" + modules = _linear_attention_modules(model) + # getattr also handles pickled configs that predate the policy field. + for entry in getattr(config, "linear_attention", []): + matches = [name for name in modules if fnmatch.fnmatch(name, entry.module_name)] + matched_on_rank = [bool(matches)] + if dist.size() > 1: + # Shared recipes can select layers owned by another pipeline stage. + matched_on_rank = [False] * dist.size() + torch.distributed.all_gather_object(matched_on_rank, bool(matches)) + if not any(matched_on_rank): + raise ValueError( + f"linear_attention rule {entry.module_name!r} matches no supported modules" + ) + for name in matches: + modules[name].linear_attention_config = entry.cfg.model_copy(deep=True) + + +def _validate_linear_attention(model): + """Validate configured modules after both policy and quantizers are assigned.""" + for module in _linear_attention_modules(model).values(): + module.validate_linear_attention() + + +def _linear_attention_state(model): + return { + name: module.linear_attention_config.model_dump() + for name, module in _linear_attention_modules(model).items() + } + + +def _restore_linear_attention_policy(model, saved_policies): + """Restore saved policies; modelopt_post_restore validates the completed state.""" + if saved_policies is None: + return + modules = _linear_attention_modules(model) + if saved_policies.keys() != modules.keys(): + raise ApplyModeError("Saved linear_attention policies do not match the restored modules") + for name, policy in saved_policies.items(): + modules[name].linear_attention_config = LinearAttentionConfig(**policy) + + +def _restore_legacy_linear_attention_quantizers(model, quantizer_state): + """Keep newly introduced, disabled handles loadable from older checkpoints.""" + for name, module in _linear_attention_modules(model).items(): + for handle in module.linear_attention_quantizer_names: + key = f"{name}.{handle}" if name else handle + quantizer = getattr(module, handle) + if key not in quantizer_state and not quantizer.is_enabled: + quantizer_state[key] = quantizer.get_modelopt_state() + + +class _LinearAttentionQuantMixin(QuantModule): + linear_attention_quantizer_names: tuple[str, ...] = () + + def _setup(self): + for name in self.linear_attention_quantizer_names: + self._register_temp_attribute( + name, TensorQuantizer(QuantizerAttributeConfig(enable=False)) + ) + self._register_temp_attribute("linear_attention_config", LinearAttentionConfig()) + self._register_temp_attribute("_linear_attention_prefill_lengths", None) + + @property + def _linear_attn_state(self): + return getattr(self, self.linear_attention_quantizer_names[0]) + + @property + def linear_attention_is_enabled(self): + """Whether state quantization or the serving arithmetic policy is enabled.""" + return ( + any(getattr(self, name).is_enabled for name in self.linear_attention_quantizer_names) + or self.linear_attention_config.backend == "serving" + ) + + def validate_linear_attention(self): + """Validate quantizer contracts shared by GDN and KDA.""" + if self._linear_attn_state.is_enabled: + state_format, group_size = state_quantizer_config( + self._linear_attn_state, + name=self.linear_attention_quantizer_names[0], + ) + if self.linear_attention_config.backend != "serving": + raise ValueError( + "GDN/KDA state QAT requires backend='serving' and prefill lengths " + "through linear_attention_training_phase" + ) + if self.linear_attention_config.state_codec == "int8_hadamard32": + if state_format != "int8": + raise ValueError("int8_hadamard32 requires INT8 state quantization") + if group_size: + raise ValueError("TensorQuantizer block_sizes requires state_codec='tile'") + + def modelopt_post_restore(self, prefix=""): + """Validate the restored numerical policy and quantizers.""" + super().modelopt_post_restore(prefix) + self.validate_linear_attention() diff --git a/modelopt/torch/quantization/plugins/megatron.py b/modelopt/torch/quantization/plugins/megatron.py index 5e422235e14..63c6acc029e 100644 --- a/modelopt/torch/quantization/plugins/megatron.py +++ b/modelopt/torch/quantization/plugins/megatron.py @@ -58,6 +58,7 @@ from ..algorithms import AutoQuantizeGradientSearcher from ..conversion import maybe_promote_nvfp4_static_quantizer +from ..linear_attention.config import LinearAttentionConfig from ..nn import ( GroupedQuantizer, QuantModule, @@ -70,7 +71,9 @@ from ..utils import sync_moe_expert_amax from ..utils.layerwise_calib import LayerActivationCollector from .custom import CUSTOM_MODEL_PLUGINS, _ParallelLinear -from .gated_delta_net import GatedDeltaNetStateQuantMixin +from .gdn import GatedDeltaNetStateQuantMixin +from .kda import KimiDeltaAttentionStateQuantMixin +from .linear_attention import _LinearAttentionQuantMixin from .transformer_engine import _QuantTEGroupedLinear, _QuantTELayerNormLinear, _QuantTELinear try: @@ -87,6 +90,13 @@ except ImportError: HAS_GDN = False +try: + from megatron.core.ssm.gated_delta_net import KimiDeltaAttention + + HAS_KDA = True +except ImportError: + HAS_KDA = False + __all__ = [] @@ -210,6 +220,8 @@ def quant_module_get_extra_state(self) -> dict: quantizer_state[name] = module.get_modelopt_state() extra_state["modelopt_quantizer_state"] = quantizer_state + if isinstance(self, _LinearAttentionQuantMixin): + extra_state["modelopt_linear_attention_state"] = self.linear_attention_config.model_dump() # Handle real_quantizer_state and q_tensor_state extra_state.update(real_quant_module_get_extra_state(self)) @@ -288,6 +300,12 @@ def quant_module_set_extra_state(self, state: Any): if state is None or not self.allow_post_restore: return + if isinstance(self, _LinearAttentionQuantMixin) and "modelopt_linear_attention_state" in state: + # Restore the policy before quantizer restoration invokes modelopt_post_restore. + self.linear_attention_config = LinearAttentionConfig( + **state["modelopt_linear_attention_state"] + ) + quantizer_state = state.get("modelopt_quantizer_state", None) if quantizer_state is not None: @@ -1106,63 +1124,142 @@ def set_extra_state(self, state): quant_module_set_extra_state(self, state) -if HAS_GDN: +class _MegatronLinearAttentionMixin(_LinearAttentionQuantMixin): + """Wrap the kernel call in both direct and recomputed Megatron forwards.""" - @QuantModuleRegistry.register({GatedDeltaNet: "megatron_GatedDeltaNet"}) - class _QuantGatedDeltaNet(GatedDeltaNetStateQuantMixin): - """GatedDeltaNet with fake quantization of the recurrent state at kernel chunk boundaries. - - Routes ``self.gated_delta_rule`` through state/W QDQ from both Megatron's direct - forward and older split-forward layouts. The older dynamic-batching inference - paths (``ssm_prefill`` / ``ssm_decode``) are left untouched. - """ + # PyTorch requires class-level overrides to save/load ``_extra_state``; + # instance-level ModelOpt callbacks alone are not sufficient for GDN or KDA. + def get_extra_state(self): + return quant_module_get_extra_state(self) - # Class-level overrides so torch routes the quantizer state through ``_extra_state`` - # (GatedDeltaNet has none); see _QuantDSAttention. - def get_extra_state(self): - return quant_module_get_extra_state(self) + def set_extra_state(self, state): + quant_module_set_extra_state(self, state) - def set_extra_state(self, state): - quant_module_set_extra_state(self, state) + def _setup(self): + super()._setup() + self._register_temp_attribute("_linear_attention_replay_gate_inputs", None) + try: + data_parallel_group = get_data_parallel_group(with_context_parallel=True) + except AssertionError: + data_parallel_group = get_data_parallel_group() + self.parallel_state = ParallelState( + data_parallel_group, + mcore_parallel.get_tensor_model_parallel_group(), + ) - def _setup(self): - super()._setup() - try: - data_parallel_group = get_data_parallel_group(with_context_parallel=True) - except AssertionError: - data_parallel_group = get_data_parallel_group() - self.parallel_state = ParallelState( - data_parallel_group, - mcore_parallel.get_tensor_model_parallel_group(), + @property + def _serving_arithmetic(self): + return self.linear_attention_config.backend == "serving" + + def _prepare_input_for_gated_delta_rule(self, *args, **kwargs): + # Preserve raw BF16 Q/K: prefill stores normalized BF16 operands, whereas + # decode normalizes inside its FP32 update. One shared pre-normalization + # would irreversibly change the suffix state trajectory. + normalize = self.use_qk_l2norm + if self._serving_arithmetic: + self.use_qk_l2norm = False + try: + return super()._prepare_input_for_gated_delta_rule(*args, **kwargs) + finally: + self.use_qk_l2norm = normalize + + @contextmanager + def _quantized_linear_attention_kernel(self): + kernel = self.gated_delta_rule + previous_gates = self._linear_attention_replay_gate_inputs + if self._serving_arithmetic: + + def run_kernel(*args, **kwargs): + kwargs["use_qk_l2norm_in_kernel"] = self.use_qk_l2norm + if self._linear_attention_replay_gate_inputs is not None: + kwargs["replay_gate_inputs"] = self._linear_attention_replay_gate_inputs + return self._linear_attention_kernel(kernel, *args, **kwargs) + + self.gated_delta_rule = run_kernel + else: + self.gated_delta_rule = partial(self._linear_attention_kernel, kernel) + self._linear_attention_replay_gate_inputs = None + try: + yield + finally: + self.gated_delta_rule = kernel + self._linear_attention_replay_gate_inputs = previous_gates + + def forward(self, *args, **kwargs): + # Newer Megatron versions recompute the core independently during backward. + if hasattr(super(), "_forward_compute") or hasattr( + super(), "forward_pre_attn_and_core_attn" + ): + return super().forward(*args, **kwargs) + with self._quantized_linear_attention_kernel(): + return super().forward(*args, **kwargs) + + def _forward_compute(self, *args, **kwargs): + with self._quantized_linear_attention_kernel(): + return super()._forward_compute(*args, **kwargs) + + def forward_pre_attn_and_core_attn(self, *args, **kwargs): + with self._quantized_linear_attention_kernel(): + return super().forward_pre_attn_and_core_attn(*args, **kwargs) + + def validate_linear_attention(self): + super().validate_linear_attention() + if self.config.context_parallel_size > 1 and self.linear_attention_is_enabled: + raise NotImplementedError("GDN/KDA QAT does not support Megatron context parallelism.") + if self._serving_arithmetic and ( + not hasattr(super(), "_prepare_input_for_gated_delta_rule") + or getattr(self, "gdn_pre_gated_delta_rule_fusion", False) + ): + raise NotImplementedError( + "Serving arithmetic requires Megatron's unfused input-preparation hook" ) - @contextmanager - def _quantized_gdn_kernel(self): - gated_delta_rule = self.gated_delta_rule - self.gated_delta_rule = partial( - self._state_quantized_chunk_gated_delta_rule, gated_delta_rule - ) - try: - yield - finally: - self.gated_delta_rule = gated_delta_rule - - def forward(self, *args, **kwargs): - if hasattr(GatedDeltaNet, "forward_pre_attn_and_core_attn"): - return super().forward(*args, **kwargs) - with self._quantized_gdn_kernel(): - return super().forward(*args, **kwargs) - - def forward_pre_attn_and_core_attn(self, *args, **kwargs): - with self._quantized_gdn_kernel(): - return super().forward_pre_attn_and_core_attn(*args, **kwargs) - - def validate_linear_attention(self): - super().validate_linear_attention() - if self.config.context_parallel_size > 1 and ( - self.gdn_state_quantizer.is_enabled or self.gdn_w_quantizer.is_enabled - ): - raise NotImplementedError("GDN QAT does not support Megatron context parallelism.") + +if HAS_GDN: + + @QuantModuleRegistry.register({GatedDeltaNet: "megatron_GatedDeltaNet"}) + class _QuantGatedDeltaNet(_MegatronLinearAttentionMixin, GatedDeltaNetStateQuantMixin): + """Megatron GDN with serving-aligned state fake quantization.""" + + _linear_attention_kernel = ( + GatedDeltaNetStateQuantMixin._state_quantized_chunk_gated_delta_rule + ) + + def _compute_gates(self, a_log, dt_bias, batch, seq_len, *gate_feats): + gate, inputs = super()._compute_gates(a_log, dt_bias, batch, seq_len, *gate_feats) + if self._serving_arithmetic: + # Import the optional vLLM backend only for the native precision profile. + from ...kernels.quantization.linear_attention.serving.forward import ( + fused_gdn_gating, + ) + from ..linear_attention.utils import forward_value + + raw_beta, raw_gate = gate_feats + if self.linear_attention_config.precision == "replayssm": + self._linear_attention_replay_gate_inputs = (raw_gate, raw_beta, a_log, dt_bias) + with torch.no_grad(): + native_gate, native_beta = fused_gdn_gating( + a_log, + raw_gate.reshape(-1, raw_gate.shape[-1]).contiguous(), + raw_beta.reshape(-1, raw_beta.shape[-1]).contiguous(), + dt_bias, + ) + gate = forward_value(gate, native_gate.reshape_as(gate)) + inputs["beta"] = forward_value( + inputs["beta"], native_beta.reshape_as(inputs["beta"]) + ).to(raw_beta.dtype) + return gate, inputs + + +if HAS_KDA: + + @QuantModuleRegistry.register({KimiDeltaAttention: "megatron_KimiDeltaAttention"}) + class _QuantKimiDeltaAttention( + _MegatronLinearAttentionMixin, KimiDeltaAttentionStateQuantMixin + ): + """Megatron KDA with decode-aware state fake quantization.""" + + _linear_attention_kernel = KimiDeltaAttentionStateQuantMixin._state_quantized_chunk_kda def _is_supported_megatron_model(model: torch.nn.Module) -> bool: diff --git a/modelopt/torch/quantization/plugins/vllm_linear_attention.py b/modelopt/torch/quantization/plugins/vllm_linear_attention.py new file mode 100644 index 00000000000..49131c9a481 --- /dev/null +++ b/modelopt/torch/quantization/plugins/vllm_linear_attention.py @@ -0,0 +1,176 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Recurrent-state QDQ at native vLLM prefill and decode call boundaries.""" + +from functools import partial +from types import FunctionType + +import torch +import vllm +from packaging.version import Version +from vllm.forward_context import get_forward_context + +from ..linear_attention.utils import _state_qdq, state_quantizer_config +from ..nn import QuantModuleRegistry +from .custom import CUSTOM_MODEL_PLUGINS +from .linear_attention import _LinearAttentionQuantMixin + +__all__ = ["bind_vllm_linear_attention"] + + +def bind_vllm_linear_attention(model, model_runner): + """Validate state-only adapters before calibration or generation.""" + layers = [m for m in model.modules() if isinstance(m, _QuantVllmLinearAttention)] + for layer in layers: + layer.validate_linear_attention() + if any(layer.linear_attention_is_enabled for layer in layers): + _validate_runtime(model_runner) + return layers + + +def _validate_runtime(runner): + version = Version(vllm.__version__) + if version.release[:2] != (0, 15): + raise NotImplementedError("Linear-attention fakequant currently requires vLLM 0.15.x V1") + from vllm.model_executor.layers.mamba.mamba_utils import MambaStateShapeCalculator + + shape = MambaStateShapeCalculator.gated_delta_net_state_shape(1, 2, 2, 32, 64, 4)[-1] + if shape[-2:] != (32, 64): + raise NotImplementedError( + "Linear-attention fakequant requires the key-first vLLM state ABI" + ) + config = runner.vllm_config + if not config.model_config.enforce_eager or config.scheduler_config.async_scheduling: + raise ValueError("State fakequant requires --enforce-eager --no-async-scheduling") + if config.speculative_config is not None or config.cache_config.enable_prefix_caching: + raise ValueError("Disable speculative decoding and prefix caching for state fakequant") + if config.cache_config.mamba_cache_mode != "none": + raise ValueError("State fakequant requires mamba_cache_mode='none'") + if config.kv_transfer_config is not None or config.ec_transfer_config is not None: + raise ValueError("State fakequant does not support state transfer") + parallel = config.parallel_config + if any( + getattr(parallel, name, 1) != 1 + for name in ( + "pipeline_parallel_size", + "data_parallel_size", + "prefill_context_parallel_size", + "decode_context_parallel_size", + ) + ): + raise ValueError("State fakequant supports TP with PP=1, DP=1, and CP=1") + + +class _QuantVllmLinearAttention(_LinearAttentionQuantMixin): + linear_attention_quantizer_names = ("gdn_state_quantizer", "gdn_w_quantizer") + + def validate_linear_attention(self): + # vLLM owns phase boundaries; the training-only validation needs a context. + if self._linear_attn_state.is_enabled: + state_quantizer_config(self._linear_attn_state) + if ( + any( + getattr(self, name).is_enabled for name in self.linear_attention_quantizer_names[1:] + ) + or self.linear_attention_config.precision == "replayssm" + ): + raise ValueError( + "vLLM state fakequant supports plain state QDQ; W and ReplaySSM are unsupported" + ) + + def _state_qdq(self, state): + return _state_qdq( + state, + block_v=self.linear_attention_config.state_block_v, + state_format="int8" if self._linear_attn_state.num_bits == 8 else "fp8_e4m3", + state_quantizer=self._linear_attn_state, + ) + + def _quantized_state_call(self, native, *args, **kwargs): + state = kwargs.get("initial_state") + if state is not None and self._linear_attn_state.is_enabled: + if state.dtype != torch.float32: + raise ValueError("State fakequant requires an FP32 recurrent cache") + indices = kwargs.get("ssm_state_indices") + if indices is None: + # Native prefill has already initialized fresh requests to zero. + kwargs["initial_state"] = self._state_qdq(state) + else: + # Eager, non-speculative decode has one cache slot per active sequence. + indices = indices[: kwargs["cu_seqlens"].numel() - 1].long() + state.index_copy_(0, indices, self._state_qdq(state.index_select(0, indices))) + return native(*args, **kwargs) + + def _forward_with_state_qdq(self, original, kernel_names, *args, **kwargs): + if not self.linear_attention_is_enabled: + return original(*args, **kwargs) + context = get_forward_context() + if context.attn_metadata is None: + return original(*args, **kwargs) + if context.attn_metadata[self.prefix].spec_sequence_masks is not None: + raise ValueError("Speculative metadata is unsupported for state fakequant") + function = original.__func__ + replacements = { + name: partial(self._quantized_state_call, function.__globals__[name]) + for name in kernel_names + } + # Bind wrappers for this invocation only; every wrapper calls the original kernel. + forward = FunctionType( + function.__code__, + {**function.__globals__, **replacements}, + function.__name__, + function.__defaults__, + function.__closure__, + ) + forward.__kwdefaults__ = function.__kwdefaults__ + return forward(self, *args, **kwargs) + + +class _QuantVllmGDN(_QuantVllmLinearAttention): + def _forward_core(self, *args, **kwargs): + return self._forward_with_state_qdq( + super()._forward_core, + ("chunk_gated_delta_rule", "fused_recurrent_gated_delta_rule"), + *args, + **kwargs, + ) + + +class _QuantVllmKDA(_QuantVllmLinearAttention): + linear_attention_quantizer_names = ("kda_state_quantizer",) + + def _forward(self, *args, **kwargs): + return self._forward_with_state_qdq( + super()._forward, + ("chunk_kda", "fused_recurrent_kda"), + *args, + **kwargs, + ) + + +def _register_vllm_linear_attention(model): + adapters = { + ("vllm.model_executor.models.qwen3_next", "Qwen3NextGatedDeltaNet"): _QuantVllmGDN, + ("vllm.model_executor.layers.kda", "KimiDeltaAttention"): _QuantVllmKDA, + } + for module in model.modules(): + cls = type(module) + adapter = adapters.get((cls.__module__, cls.__name__)) + if adapter is not None and cls not in QuantModuleRegistry: + QuantModuleRegistry.register({cls: f"vllm_{cls.__name__}"})(adapter) + + +CUSTOM_MODEL_PLUGINS.add(_register_vllm_linear_attention) diff --git a/modelopt_recipes/configs/ptq/units/README.md b/modelopt_recipes/configs/ptq/units/README.md index f98c6f01051..09098eb54ee 100644 --- a/modelopt_recipes/configs/ptq/units/README.md +++ b/modelopt_recipes/configs/ptq/units/README.md @@ -33,7 +33,10 @@ recipes (under `general/` or `models/`) or presets (under `presets/`). | `experts_nvfp4.yaml` | NVFP4 W4A4 on `*.experts.*` weight/input quantizers | | `mixer_mlp_nvfp4.yaml` | NVFP4 W4A4 on dense `*.mixer.{up,down}_proj` weight/input quantizers | | `attention_qkv_fp8.yaml` | FP8 E4M3 on attention q/k/v bmm and softmax quantizers | -| `gdn_state_fp8_dynamic.yaml` | FP8 E4M3 dynamic fake quantization of the GatedDeltaNet recurrent state (per sequence, head and 64-column tile) at every kernel chunk boundary; needs `fla-core==0.5.1`, Triton, and SM89+ | -| `gdn_w_fp8_dynamic.yaml` | FP8 E4M3 dynamic (per token and head) fake quantization of the WY tensor `w` that multiplies the GatedDeltaNet state; requires identity STE, `fla-core==0.5.1`, and Triton | +| `gdn_state_fp8_dynamic.yaml` | FP8 E4M3 dynamic fake quantization of the GatedDeltaNet recurrent state (per sequence, head and 64-column tile) at serving cache boundaries; requires an explicit serving execution policy | +| `linear_attention_state_int8_dynamic.yaml` | GDN/KDA INT8 state quantizer entries; the [complete recipe](../../../general/ptq/linear_attention_state_int8_dynamic.yaml) enables Hadamard during decode | +| `linear_attention_state_int8_block32_dynamic.yaml` | Standard dynamic INT8 with one scale per key row and 32 value channels; the [complete training recipe](../../../general/ptq/linear_attention_state_int8_block32_dynamic.yaml) uses token QDQ and working-state readout | | `indexer_k_nvfp4.yaml` | NVFP4 fake quantization of the sparse-attention indexer key cache (`*indexer_k_quantizer`), global scale fixed to 1; Blackwell+ GPUs | | `indexer_q_nvfp4.yaml` | NVFP4 fake quantization of the sparse-attention indexer query (`*indexer_q_quantizer`), global scale fixed to 1; Blackwell+ GPUs | + +Native ReplaySSM uses BF16 key/update vectors; they are not independently quantized. diff --git a/modelopt_recipes/configs/ptq/units/default_disabled_quantizers.yaml b/modelopt_recipes/configs/ptq/units/default_disabled_quantizers.yaml index cac4f6d80aa..8618024222f 100644 --- a/modelopt_recipes/configs/ptq/units/default_disabled_quantizers.yaml +++ b/modelopt_recipes/configs/ptq/units/default_disabled_quantizers.yaml @@ -21,6 +21,8 @@ enable: false - quantizer_name: '*gdn_w_quantizer' enable: false + - quantizer_name: '*kda_state_quantizer' + enable: false - quantizer_name: '*block_sparse_moe.gate*' enable: false - quantizer_name: '*linear_attn.conv1d*' diff --git a/modelopt_recipes/configs/ptq/units/gdn_state_fp8_dynamic.yaml b/modelopt_recipes/configs/ptq/units/gdn_state_fp8_dynamic.yaml index ed5768ee599..ac36caab0f4 100644 --- a/modelopt_recipes/configs/ptq/units/gdn_state_fp8_dynamic.yaml +++ b/modelopt_recipes/configs/ptq/units/gdn_state_fp8_dynamic.yaml @@ -14,12 +14,15 @@ # limitations under the License. # Dynamic per-tile FP8 E4M3 fake quantization of the GatedDeltaNet recurrent state. -# Recompute amax over each [K, 64] tile inside the kernel, independently per sequence and head -# of the [N, H, K, V] state (no calibration). ``axis: [0, 1]`` specifies sequence/head grouping; -# the kernel uses one scale per full-K by 64-value-column tile (two per head for V = 128). -# This is the only configuration the module supports. Applied at every kernel chunk boundary -# (64 tokens) during calibration, QAT and QAD; needs fla-core==0.5.1 and Triton. -# See ``modelopt.torch.quantization.plugins.gated_delta_net``. +# TensorQuantizer computes amax over each [K, BV] tile of the [N, H, K, V] state, +# independently per sequence/head (axis: [0, 1]), without calibration. +# Set BV through linear_attention[].cfg.state_block_v (default: 64). +# This unit enables only the quantizer. Select backend='serving' in the execution +# policy and provide prefill lengths through linear_attention_training_phase. +# With precision='vllm', state QDQ applies at the handoff to a nonempty decode +# suffix and at each suffix token's cache write; fresh prefill has no internal QDQ. +# The serving profile requires public vLLM's GDN kernels. +# See modelopt.torch.quantization.plugins.gdn. # modelopt-schema: modelopt.torch.quantization.config.QuantizerCfgListConfig - quantizer_name: '*gdn_state_quantizer' diff --git a/modelopt_recipes/configs/ptq/units/gdn_w_fp8_dynamic.yaml b/modelopt_recipes/configs/ptq/units/linear_attention_state_int8_block32_dynamic.yaml similarity index 57% rename from modelopt_recipes/configs/ptq/units/gdn_w_fp8_dynamic.yaml rename to modelopt_recipes/configs/ptq/units/linear_attention_state_int8_block32_dynamic.yaml index 2c73fdd6eaa..00284f84ce3 100644 --- a/modelopt_recipes/configs/ptq/units/gdn_w_fp8_dynamic.yaml +++ b/modelopt_recipes/configs/ptq/units/linear_attention_state_int8_block32_dynamic.yaml @@ -13,16 +13,21 @@ # See the License for the specific language governing permissions and # limitations under the License. -# FP8 E4M3 fake quantization of ``w``, the WY-transformed keys that multiply the GatedDeltaNet -# recurrent state inside the chunked kernel, with a dynamic scale per token and head -# (``w`` is ``[B, T, H, K]``; ``axis: [0, 1, 2]`` reduces over K). Pairs with -# ``gdn_state_fp8_dynamic`` to emulate an FP8 x FP8 state matmul. Unlike the state quantizer this -# one runs on the tensor itself. The initial integration supports dynamic E4M3 with identity -# STE only. Needs fla-core==0.5.1 and Triton. - +# Dynamic INT8 with one scale per key row and 32 value channels for training and serving. # modelopt-schema: modelopt.torch.quantization.config.QuantizerCfgListConfig - - quantizer_name: '*gdn_w_quantizer' + - quantizer_name: '*gdn_state_quantizer' + cfg: + num_bits: 8 + unsigned: false + narrow_range: true + block_sizes: {-1: 32} + type: dynamic + pass_through_bwd: true + - quantizer_name: '*kda_state_quantizer' cfg: - num_bits: e4m3 - axis: [0, 1, 2] + num_bits: 8 + unsigned: false + narrow_range: true + block_sizes: {-1: 32} type: dynamic + pass_through_bwd: true diff --git a/modelopt_recipes/configs/ptq/units/linear_attention_state_int8_dynamic.yaml b/modelopt_recipes/configs/ptq/units/linear_attention_state_int8_dynamic.yaml new file mode 100644 index 00000000000..eceadf5d5d0 --- /dev/null +++ b/modelopt_recipes/configs/ptq/units/linear_attention_state_int8_dynamic.yaml @@ -0,0 +1,34 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Quantizer-only INT8 unit for GDN and KDA recurrent states. +# The complete general/ptq/linear_attention_state_int8_dynamic recipe enables Hadamard. +# modelopt-schema: modelopt.torch.quantization.config.QuantizerCfgListConfig + - quantizer_name: '*gdn_state_quantizer' + cfg: + num_bits: 8 + unsigned: false + narrow_range: true + axis: [0, 1] + type: dynamic + pass_through_bwd: true + - quantizer_name: '*kda_state_quantizer' + cfg: + num_bits: 8 + unsigned: false + narrow_range: true + axis: [0, 1] + type: dynamic + pass_through_bwd: true diff --git a/modelopt_recipes/general/ptq/linear_attention_state_int8_block32_dynamic.yaml b/modelopt_recipes/general/ptq/linear_attention_state_int8_block32_dynamic.yaml new file mode 100644 index 00000000000..540f48a978a --- /dev/null +++ b/modelopt_recipes/general/ptq/linear_attention_state_int8_block32_dynamic.yaml @@ -0,0 +1,36 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# modelopt-schema: modelopt.recipe.config.ModelOptPTQRecipe +imports: + base_disable_all: configs/ptq/units/base_disable_all + state_int8: configs/ptq/units/linear_attention_state_int8_block32_dynamic + +metadata: + description: >- + Dynamic INT8 GDN/KDA state QDQ with per-key 32-value scales and working-state readout. + Emulates vLLM 0.15 BF16 prefill arithmetic without additional state QDQ. + Supply per-sequence prefix lengths through + linear_attention_training_phase. No calibration is required. +quantize: + algorithm: + quant_cfg: + - $import: base_disable_all + - $import: state_int8 + linear_attention: + - module_name: '*' + cfg: + backend: serving + precision: vllm diff --git a/modelopt_recipes/general/ptq/linear_attention_state_int8_dynamic.yaml b/modelopt_recipes/general/ptq/linear_attention_state_int8_dynamic.yaml new file mode 100644 index 00000000000..b21cc4c2653 --- /dev/null +++ b/modelopt_recipes/general/ptq/linear_attention_state_int8_dynamic.yaml @@ -0,0 +1,36 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# modelopt-schema: modelopt.recipe.config.ModelOptPTQRecipe +imports: + base_disable_all: configs/ptq/units/base_disable_all + state_int8: configs/ptq/units/linear_attention_state_int8_dynamic + +metadata: + description: >- + Dynamic INT8 GDN/KDA state QDQ with 32-value Hadamard rotation during decode. + Imports the quantized-ReplaySSM serving kernels and uses working-state readout. + Fresh prefill has no internal state QDQ; supply per-sequence prefix lengths through + linear_attention_training_phase. No calibration is required. +quantize: + algorithm: + quant_cfg: + - $import: base_disable_all + - $import: state_int8 + linear_attention: + - module_name: '*' + cfg: + backend: serving + precision: replayssm diff --git a/modelopt_recipes/ptq.md b/modelopt_recipes/ptq.md index 21e5fe5744b..62533f64111 100644 --- a/modelopt_recipes/ptq.md +++ b/modelopt_recipes/ptq.md @@ -28,7 +28,7 @@ supported combinations. ### The shipped recipes
-All 32 general/ptq/ recipes (click to expand) +All 34 general/ptq/ recipes (click to expand) | Recipe | Model body | KV cache | Calibration | |--------|-----------|----------|-------------| @@ -64,9 +64,14 @@ supported combinations. | `iq2_xs` | IQ2_XS W2A16 (2.31 bpw), MLP + MoE weights only | none | GPTQ (layerwise) | | `iq2_s` | IQ2_S W2A16 (2.56 bpw), MLP + MoE weights only | none | GPTQ (layerwise) | | `q8_0` | Q8_0 W8A16 (8.5 bpw), eligible linears | none | none (no calibration) | +| `linear_attention_state_int8_dynamic` | GDN/KDA decode state INT8 + Hadamard; weights unchanged | none | none (dynamic scales; requires a prefix/decode phase context) | +| `linear_attention_state_int8_block32_dynamic` | GDN/KDA state INT8, 32 value channels per key row; weights unchanged | none | none (dynamic scales; requires vLLM and a prefix/decode phase context) |
+The block32 linear-attention recipe selects `precision="vllm"` +and working-state readout. It applies fake QDQ with floating-point state storage. + --- ### Model-body schemes diff --git a/pyproject.toml b/pyproject.toml index cc95e6bca7c..b792592022b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -329,15 +329,6 @@ disable_error_code = ["attr-defined"] module = ["examples.diffusers.fastgen.preprocess.*"] ignore_errors = true -# Vendored from fla-org/flash-linear-attention (MIT); kept faithful to upstream rather than -# annotated to modelopt's strict mypy (Triton loop indices, ``*tensor.shape`` unpacking). -[[tool.mypy.overrides]] -module = [ - "modelopt.torch.kernels.quantization.linear_attention.fla_chunk_delta_h", - "modelopt.torch.kernels.quantization.linear_attention.fla_chunk_gated_delta_rule", -] -ignore_errors = true - [tool.bandit] exclude_dirs = [".github/", "examples/", "noxfile.py", "tests/"] # Do not change `skips`. It should be consistent with NVIDIA's Wheel-CI-CD bandit.yml config. diff --git a/tests/_test_utils/torch/quantization/linear_attention_reference.py b/tests/_test_utils/torch/quantization/linear_attention_reference.py index 65ea183a318..1373b5ea46b 100644 --- a/tests/_test_utils/torch/quantization/linear_attention_reference.py +++ b/tests/_test_utils/torch/quantization/linear_attention_reference.py @@ -20,28 +20,40 @@ are read on the CPU. Accumulation follows the input dtype (use float64 for algebra tests). """ -from collections.abc import Callable from itertools import pairwise import torch -__all__ = ["chunk_gdn_reference", "recurrent_delta_rule_reference", "state_fp8_qdq_reference"] +__all__ = [ + "chunk_delta_rule_reference", + "recurrent_delta_rule_reference", + "state_qdq_reference", +] -def state_fp8_qdq_reference(state: torch.Tensor, block_v: int = 64) -> torch.Tensor: - """Dynamic E4M3 QDQ with identity STE and one scale per ``[Dk, block_v]`` tile. - - ``state`` is in key-first layout ``[..., Dk, Dv]``. Scales and rounding do not - contribute derivatives. The zero tile uses scale one; partial value tiles are valid. - """ +def state_qdq_reference(state: torch.Tensor, block_v: int = 64, state_format: str = "fp8_e4m3"): + """Dynamic state-tile QDQ with detached scales and identity STE.""" + if state_format not in ("fp8_e4m3", "int8"): + raise ValueError("State format must be fp8_e4m3 or int8") if block_v not in (16, 32, 64, 128): raise ValueError("block_v must be 16, 32, 64, or 128") with torch.no_grad(): rounded = [] for tile in state.float().split(block_v, dim=-1): amax = tile.abs().amax(dim=(-2, -1), keepdim=True) - scale = torch.where(amax > 0, amax / 448.0, torch.ones_like(amax)) - rounded.append((tile / scale).clamp(-448, 448).to(torch.float8_e4m3fn).float() * scale) + if state_format == "int8": + # CUDA TensorQuantizer uses a quantization multiplier and zeros tiny groups. + tiny = amax < 2**-24 + quant_scale = 127.0 / torch.where(tiny, torch.ones_like(amax), amax) + codes = (tile * quant_scale).round().clamp(-127, 127) + decoded = torch.where(tiny, 0.0, codes / quant_scale) + else: + safe_amax = torch.where(amax <= 2**-24, torch.ones_like(amax), amax) + quant_scale = torch.div(448.0, safe_amax) + scale = quant_scale.reciprocal() + codes = (tile * quant_scale).clamp(-448, 448).to(torch.float8_e4m3fn).float() + decoded = codes * scale + rounded.append(decoded) quantized = torch.cat(rounded, dim=-1).to(state.dtype) return state + (quantized - state).detach() @@ -126,7 +138,7 @@ def recurrent_delta_rule_reference( return output, final.transpose(-1, -2) if state_v_first else final -def chunk_gdn_reference( +def chunk_gdn( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, @@ -135,69 +147,114 @@ def chunk_gdn_reference( *, chunk_size: int = 64, scale: float | None = None, - initial_state: torch.Tensor | None = None, - cu_seqlens: torch.Tensor | None = None, - state_v_first: bool = False, + initial_state: torch.Tensor, state_qdq: bool = False, state_qdq_block_v: int = 64, - w_quantizer: Callable[[torch.Tensor], torch.Tensor] | None = None, + state_format: str = "fp8_e4m3", ) -> tuple[torch.Tensor, torch.Tensor]: - """Exact GDN chunk algebra with optional state/W fake quantization. - - The solve is a unit-lower triangular solve. ``w_quantizer`` sees the complete - materialized ``[B,T,Hv,Dk]`` WY operand once, with its own autograd semantics. - State QDQ occurs on the initial state and each chunk's final state, after readout. - """ - if g.ndim != 3 or chunk_size <= 0: - raise ValueError("chunk_gdn_reference requires scalar GDN gates and positive chunk_size") - q, k, states, sequences = _prepare(q, k, v, g, beta, initial_state, cu_seqlens, state_v_first) + """Compute one prepared, nonempty GDN prefix [T,H,D] with a key-first state.""" scale = q.shape[-1] ** -0.5 if scale is None else scale - chunks, all_w = [], [] - for n, (b, start, end) in enumerate(sequences): - for lo in range(start, end, chunk_size): - hi = min(lo + chunk_size, end) - qc, kc, vc = (x[b, lo:hi].transpose(0, 1) for x in (q, k, v)) - gc = g[b, lo:hi].transpose(0, 1).cumsum(-1) - bc = beta[b, lo:hi].transpose(0, 1).unsqueeze(-1) - # Mask before exp: upper-triangle positive differences can overflow for long decay. - causal = torch.ones(hi - lo, hi - lo, device=q.device, dtype=torch.bool).tril() - decay = (gc.unsqueeze(-1) - gc.unsqueeze(-2)).masked_fill(~causal, 0).exp() - gram = kc @ kc.transpose(-1, -2) - lower = (bc * gram * decay).tril(-1) - matrix = lower + torch.eye(hi - lo, device=q.device, dtype=q.dtype) - rhs = torch.cat((bc * vc, bc * kc * gc.exp().unsqueeze(-1)), dim=-1) - solved = torch.linalg.solve_triangular(matrix, rhs, upper=False, unitriangular=True) - u, w = solved.split((v.shape[-1], k.shape[-1]), dim=-1) - 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: - w = w_quantizer(w) - w = w.reshape(-1, *w.shape[2:]) - outputs, finals, offset, previous_n = [], [], 0, -1 - for n, qc, kc, gc, decay, u in chunks: - if n != previous_n: - state = states[n] - if state_qdq: - state = state_fp8_qdq_reference(state, state_qdq_block_v) - previous_n = n - length = qc.shape[1] - wc = w[offset : offset + length].transpose(0, 1) - offset += length - updated_values = u - wc @ state # state_read - local_scores = ((qc * scale) @ kc.transpose(-1, -2) * decay).tril() + state = initial_state + if state_qdq: + state = state_qdq_reference(state, state_qdq_block_v, state_format) + pieces = [] + for lo in range(0, len(q), chunk_size): + hi = min(lo + chunk_size, len(q)) + qc, kc, vc = (x[lo:hi].transpose(0, 1) for x in (q, k, v)) + gc = g[lo:hi].transpose(0, 1).cumsum(-1) + bc = beta[lo:hi].transpose(0, 1).unsqueeze(-1) + # Mask before exp: upper-triangle positive differences can overflow for long decay. + causal = torch.ones(hi - lo, hi - lo, device=q.device, dtype=torch.bool).tril() + decay = (gc.unsqueeze(-1) - gc.unsqueeze(-2)).masked_fill(~causal, 0).exp() + lower = (bc * (kc @ kc.transpose(-1, -2)) * decay).tril(-1) + matrix = lower + torch.eye(hi - lo, device=q.device, dtype=q.dtype) + gate = gc.exp().unsqueeze(-1) + rhs = torch.cat((bc * vc, bc * kc * gate), dim=-1) + solved = torch.linalg.solve_triangular(matrix, rhs, upper=False, unitriangular=True) + u, w = solved.split((v.shape[-1], k.shape[-1]), dim=-1) + updated_values = u - w @ state + local_scores = ((qc * scale @ kc.transpose(-1, -2)) * decay).tril() output = (qc * (scale * gc.exp()).unsqueeze(-1)) @ state - output = output + local_scores @ updated_values # local_readout - outputs.append(output.transpose(0, 1)) + pieces.append((output + local_scores @ updated_values).transpose(0, 1)) weighted_keys = kc * (gc[..., -1:] - gc).exp().unsqueeze(-1) state = state * gc[..., -1].exp()[:, None, None] - state = state + weighted_keys.transpose(-1, -2) @ updated_values # state_update + state = state + weighted_keys.transpose(-1, -2) @ updated_values if state_qdq: - state = state_fp8_qdq_reference(state, state_qdq_block_v) - if len(finals) <= n: - finals.append(state) - else: - finals[n] = state - output = torch.cat(outputs).reshape(*v.shape) + state = state_qdq_reference(state, state_qdq_block_v, state_format) + return torch.cat(pieces), state + + +def chunk_kda( + q, + k, + v, + g, + beta, + *, + chunk_size=64, + scale=None, + initial_state, + state_qdq=False, + state_qdq_block_v=64, + state_format="fp8_e4m3", +): + """Compute one prepared, nonempty KDA prefix [T,H,D] with a key-first state.""" + scale = q.shape[-1] ** -0.5 if scale is None else scale + state = initial_state + if state_qdq: + state = state_qdq_reference(state, state_qdq_block_v, state_format) + pieces = [] + for lo in range(0, len(q), chunk_size): + hi = min(lo + chunk_size, len(q)) + qc, kc, vc = (x[lo:hi].transpose(0, 1) for x in (q, k, v)) + prefix = g[lo:hi].transpose(0, 1).cumsum(-2) + bc = beta[lo:hi].transpose(0, 1).unsqueeze(-1) + gate = prefix.exp() + lower_rows, score_rows = [], [] + for row in range(hi - lo): + decayed_keys = ( + kc[..., : row + 1, :] + * (prefix[..., row : row + 1, :] - prefix[..., : row + 1, :]).exp() + ) + right = decayed_keys.transpose(-1, -2) + left = kc[..., row : row + 1, :] * bc[..., row : row + 1, :] + interaction = left @ right + score = qc[..., row : row + 1, :] * scale @ right + padding = hi - lo - row - 1 + lower_rows.append(torch.nn.functional.pad(interaction, (0, padding))) + score_rows.append(torch.nn.functional.pad(score, (0, padding))) + lower = torch.cat(lower_rows, dim=-2).tril(-1) + scores = torch.cat(score_rows, dim=-2) + identity = torch.eye(hi - lo, device=q.device, dtype=q.dtype).expand_as(lower) + inverse = torch.linalg.solve_triangular( + identity + lower, identity, upper=False, unitriangular=True + ) + u = inverse @ (bc * vc) + w = inverse @ (bc * kc * gate) + updated = u - w @ state + pieces.append(((qc * scale * gate) @ state + scores @ updated).transpose(0, 1)) + weighted_keys = kc * (prefix[..., -1:, :] - prefix).exp() + state = state * gate[..., -1, :, None] + weighted_keys.transpose(-1, -2) @ updated + if state_qdq: + state = state_qdq_reference(state, state_qdq_block_v, state_format) + return torch.cat(pieces), state + + +def chunk_delta_rule_reference( + q, k, v, g, beta, *, initial_state=None, cu_seqlens=None, state_v_first=False, chunk_size=64 +): + """Evaluate GDN/KDA chunk algebra independently for each packed or batched sequence.""" + q, k, states, sequences = _prepare(q, k, v, g, beta, initial_state, cu_seqlens, state_v_first) + chunk = chunk_kda if g.ndim == 4 else chunk_gdn + outputs, finals = [], [] + for n, (b, start, end) in enumerate(sequences): + output, state = chunk( + *(x[b, start:end] for x in (q, k, v, g, beta)), + initial_state=states[n], + chunk_size=chunk_size, + ) + outputs.append(output) + finals.append(state) + output = torch.stack(outputs) if cu_seqlens is None else torch.cat(outputs).unsqueeze(0) final = torch.stack(finals) return output, final.transpose(-1, -2) if state_v_first else final diff --git a/tests/examples/megatron_bridge/test_linear_attention.py b/tests/examples/megatron_bridge/test_linear_attention.py new file mode 100644 index 00000000000..10f56993414 --- /dev/null +++ b/tests/examples/megatron_bridge/test_linear_attention.py @@ -0,0 +1,146 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import runpy + +import pytest +import torch +from _test_utils.examples.megatron_example_runner import reset_megatron_global_state +from _test_utils.examples.run_command import MODELOPT_ROOT +from megatron.bridge.models.hybrid.hybrid_provider import HybridModelProvider +from megatron.bridge.training.config import ( + CheckpointConfig, + ConfigContainer, + DistributedDataParallelConfig, + LoggerConfig, + MockGPTDatasetConfig, + OptimizerConfig, + RNGConfig, + SchedulerConfig, + TokenizerConfig, + TrainingConfig, + ValidationConfig, +) +from megatron.core.utils import unwrap_model + +from modelopt.recipe import load_recipe +from modelopt.torch.quantization.utils import is_quantized + +run_training = runpy.run_path(str(MODELOPT_ROOT / "examples/llm_qat/linear_attention/train.py"))[ + "run_training" +] + + +def _train(qad, recipe): + captured = {} + + def provider(name): + model = HybridModelProvider( + num_layers=1, + hidden_size=64, + ffn_hidden_size=128, + num_attention_heads=2, + hybrid_layer_pattern="G", + vocab_size=128, + seq_length=16, + linear_num_key_heads=2, + linear_num_value_heads=2, + linear_key_head_dim=32, + linear_value_head_dim=32, + linear_conv_kernel_dim=4, + linear_attention_freq=1, + experimental_attention_variant="gated_delta_net", + is_hybrid_model=True, + activation_func=torch.nn.functional.silu, + calculate_per_token_loss=True, + gradient_accumulation_fusion=False, + recompute_granularity="full", + recompute_method="uniform", + recompute_num_layers=1, + cross_entropy_loss_fusion=False, + ) + + def capture(models): + module = unwrap_model(models[0]) + captured[name] = module + if name == "student": + captured["before"] = ( + module.decoder.layers[0].self_attention.out_proj.weight.detach().clone() + ) + return models + + model.register_post_wrap_hook(capture) + return model + + config = ConfigContainer( + model=provider("student"), + train=TrainingConfig(train_iters=1, global_batch_size=1, micro_batch_size=1), + validation=ValidationConfig(eval_iters=0, eval_interval=1), + optimizer=OptimizerConfig( + optimizer="adam", lr=1e-2, min_lr=0, weight_decay=0, use_distributed_optimizer=True + ), + scheduler=SchedulerConfig( + lr_decay_style="constant", + lr_warmup_iters=0, + start_weight_decay=0, + end_weight_decay=0, + ), + ddp=DistributedDataParallelConfig( + average_in_collective=False, use_distributed_optimizer=True + ), + dataset=MockGPTDatasetConfig( + seq_length=16, + random_seed=123, + reset_position_ids=False, + reset_attention_mask=False, + eod_mask_loss=False, + dataloader_type="single", + num_workers=0, + ), + tokenizer=TokenizerConfig(tokenizer_type="NullTokenizer", vocab_size=128), + checkpoint=CheckpointConfig(async_save=False), + logger=LoggerConfig(log_interval=1), + rng=RNGConfig(seed=123), + mixed_precision="bf16_mixed", + ) + run_training(config, recipe, 8, provider("teacher") if qad else None) + return captured + + +@pytest.fixture(scope="module") +def compiled_state_training(): + """Compile one tiny GDN shape before timing the single-GPU training checks.""" + if not torch.cuda.is_available(): + pytest.skip("Requires CUDA") + pytest.importorskip("vllm.model_executor.layers.fla.ops.kda", exc_type=ModuleNotFoundError) + recipe = load_recipe("general/ptq/linear_attention_state_int8_block32_dynamic").quantize + try: + _train(True, recipe) + finally: + reset_megatron_global_state() + return recipe + + +@pytest.mark.parametrize("qad", [False, True], ids=["qat", "qad"]) +def test_state_training(compiled_state_training, qad): + captured = _train(qad, compiled_state_training) + attention = captured["student"].decoder.layers[0].self_attention + assert attention.gdn_state_quantizer.is_enabled + assert torch.isfinite(attention.out_proj.weight).all() + assert not torch.equal(attention.out_proj.weight, captured["before"]) + if qad: + teacher = captured["teacher"] + assert not is_quantized(teacher) + assert not any(parameter.requires_grad for parameter in teacher.parameters()) diff --git a/tests/gpu/torch/kernels/quantization/linear_attention/test_fla_chunk_gated_delta_rule.py b/tests/gpu/torch/kernels/quantization/linear_attention/test_fla_chunk_gated_delta_rule.py deleted file mode 100644 index defe05aa5d5..00000000000 --- a/tests/gpu/torch/kernels/quantization/linear_attention/test_fla_chunk_gated_delta_rule.py +++ /dev/null @@ -1,91 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Minimal forward/backward checks for GDN training fake quantization.""" - -import pytest -import torch -import torch.nn.functional as F -from _test_utils.torch.quantization.linear_attention_reference import chunk_gdn_reference - -from modelopt.torch.quantization.config import QuantizerAttributeConfig -from modelopt.torch.quantization.nn import TensorQuantizer - -pytest.importorskip("fla.ops.gated_delta_rule") - -from modelopt.torch.kernels.quantization.linear_attention.fla_chunk_gated_delta_rule import ( - chunk_gated_delta_rule, -) - - -def make_inputs(): - torch.manual_seed(123) - # Two chunks exercise recurrence with one shared BF16 shape for every mode. - shape = (1, 128, 1, 32) - q, k = [F.normalize(torch.randn(shape, device="cuda"), dim=-1) for _ in range(2)] - v = torch.randn(shape, device="cuda") - g = -torch.rand(shape[:3], device="cuda") * 0.1 - beta = torch.rand_like(g) - args = [x.to(torch.bfloat16).requires_grad_() for x in (q, k, v, g, beta)] - state = (torch.randn(1, 1, 32, 32, device="cuda") * 0.1).requires_grad_() - return args, state - - -def values_and_grads(fn, args, state, **kwargs): - result = fn(*args, initial_state=state, **kwargs) - torch.manual_seed(15) - probes = [torch.randn(x.shape, device=x.device, dtype=torch.float32) for x in result] - grads = torch.autograd.grad(sum((x * p).sum() for x, p in zip(result, probes)), (*args, state)) - return result, grads - - -def compare(actual, expected, tolerance): - for a, e in zip(actual, expected): - assert torch.isfinite(a).all() - error = (a.float() - e.float()).norm() - bound = tolerance * e.float().norm().clamp_min(1e-6) - assert error <= bound, ( - f"relative L2 error {(error / e.float().norm()).item():.5g} > {tolerance}" - ) - - -@pytest.fixture(scope="module", params=["disabled", "w", "state-w"]) -def compiled_gdn_case(request): - """Compile only the selected BF16 forward/backward path, outside the test-call timer.""" - state_qdq = request.param == "state-w" - if state_qdq and torch.cuda.get_device_capability() < (8, 9): - pytest.skip("State QDQ needs native E4M3 conversion (SM89+)") - quantizer = ( - TensorQuantizer(QuantizerAttributeConfig(num_bits=(4, 3), axis=(0, 1, 2), type="dynamic")) - if request.param != "disabled" - else None - ) - args, state = make_inputs() - kwargs = {"state_qdq": state_qdq, "w_quantizer": quantizer} - values_and_grads(chunk_gated_delta_rule, args, state, output_final_state=True, **kwargs) - torch.cuda.synchronize() - return args, state, kwargs - - -def test_gdn_forward_and_backward(compiled_gdn_case): - args, state, kwargs = compiled_gdn_case - reference_args = [x.detach().float().requires_grad_() for x in args] - reference_state = state.detach().clone().requires_grad_() - expected = values_and_grads(chunk_gdn_reference, reference_args, reference_state, **kwargs) - actual = values_and_grads( - chunk_gated_delta_rule, args, state, output_final_state=True, **kwargs - ) - compare(actual[0], expected[0], 0.03) - compare(actual[1], expected[1], 0.05) diff --git a/tests/gpu_megatron/torch/quantization/plugins/test_megatron_gated_delta_net.py b/tests/gpu_megatron/torch/quantization/plugins/test_megatron_gated_delta_net.py index 0756f2d0d40..2092e8b9da0 100644 --- a/tests/gpu_megatron/torch/quantization/plugins/test_megatron_gated_delta_net.py +++ b/tests/gpu_megatron/torch/quantization/plugins/test_megatron_gated_delta_net.py @@ -13,6 +13,8 @@ # See the License for the specific language governing permissions and # limitations under the License. +from contextlib import nullcontext + import pytest import torch from _test_utils.torch.megatron.models import get_mcore_gpt_model @@ -23,18 +25,18 @@ ) import modelopt.torch.quantization as mtq +from modelopt.torch.quantization.linear_attention import ( + LinearAttentionConfig, + linear_attention_training_phase, +) +from modelopt.torch.quantization.nn import TensorQuantizer -pytest.importorskip("fla") # Megatron-Core GatedDeltaNet and the state QDQ kernel need fla +pytest.importorskip("fla") # Megatron-Core GatedDeltaNet needs FLA for its baseline kernels. GatedDeltaNet = pytest.importorskip("megatron.core.ssm.gated_delta_net").GatedDeltaNet +pytest.importorskip("vllm") -from modelopt.torch.quantization.plugins.gated_delta_net import _state_qdq_chunk_gated_delta_rule from modelopt.torch.quantization.plugins.megatron import _QuantGatedDeltaNet -try: - _state_qdq_chunk_gated_delta_rule() -except RuntimeError as e: - pytest.skip(str(e), allow_module_level=True) - SEED = 1234 @@ -56,7 +58,7 @@ def _make_model(tp_size): .cuda() .eval() ) - # Retain enough history for chunk-boundary state rounding to affect the next chunk. + # Retain enough history for recurrent-state rounding to affect later tokens. with torch.no_grad(): for module in model.modules(): if isinstance(module, GatedDeltaNet): @@ -65,30 +67,56 @@ def _make_model(tp_size): return model -def _gdn_config(sites): +def _gdn_config(): return { - "quant_cfg": [{"quantizer_name": "*", "enable": False}] - + [ + "quant_cfg": [ + {"quantizer_name": "*", "enable": False}, { - "quantizer_name": f"*gdn_{site}_quantizer", + "quantizer_name": "*gdn_state_quantizer", "cfg": { - "num_bits": (4, 3), - "axis": (0, 1) if site == "state" else (0, 1, 2), + "num_bits": 8, + "block_sizes": {-1: 32}, "type": "dynamic", + "unsigned": False, + "narrow_range": True, }, - } - for site in sites + }, ], "algorithm": None, + "linear_attention": [ + { + "module_name": "decoder.layers.*.self_attention", + "cfg": { + "backend": "serving", + "state_block_v": 16, + "precision": "vllm_0_15", + }, + } + ], } -def _test_gdn_qat_helper(rank, size, checkpoint_path, sites): +def _gdn_forward(model): + original_forward = get_forward(model) + + def forward(m): + enabled = any( + getattr(getattr(layer, "linear_attention_config", None), "backend", None) == "serving" + for layer in m.modules() + ) + with linear_attention_training_phase(m, [64, 64]) if enabled else nullcontext(): + return original_forward(m) + + return forward + + +def _test_gdn_qat_helper(rank, size, checkpoint_path): initialize_for_megatron( tensor_model_parallel_size=size, pipeline_model_parallel_size=1, seed=SEED ) model = _make_model(size) - forward = get_forward(model) + forward = _gdn_forward(model) + outputs = [] handles = [ module.register_forward_hook( @@ -102,12 +130,17 @@ def _test_gdn_qat_helper(rank, size, checkpoint_path, sites): gdn_ref = outputs.copy() outputs.clear() - mtq.quantize(model, _gdn_config(sites)) + mtq.quantize(model, _gdn_config()) gdn_modules = [m for m in model.modules() if isinstance(m, _QuantGatedDeltaNet)] assert gdn_modules, "no GatedDeltaNet layer was wrapped" for module in gdn_modules: - assert module.gdn_state_quantizer.is_enabled == ("state" in sites) - assert module.gdn_w_quantizer.is_enabled == ("w" in sites) + assert module.gdn_state_qdq_block_v == 16 + assert module.gdn_state_quantizer.is_enabled + assert not module.gdn_w_quantizer.is_enabled + # Checkpointing must save the resolved policy, including edits after conversion. + module.linear_attention_config = LinearAttentionConfig( + **{**module.linear_attention_config.model_dump(), "state_block_v": 32} + ) with torch.no_grad(): loss_quant = forward(model) @@ -120,20 +153,34 @@ def _test_gdn_qat_helper(rank, size, checkpoint_path, sites): for handle in handles: handle.remove() - for site in sites: - mtq.disable_quantizer(model, f"*gdn_{site}_quantizer") + enabled = [ + q + for m in gdn_modules + for q in m.modules() + if isinstance(q, TensorQuantizer) and q.is_enabled + ] + policies = [m.linear_attention_config for m in gdn_modules] + for q in enabled: + q.disable() + for module in gdn_modules: + module.linear_attention_config = LinearAttentionConfig() with torch.no_grad(): torch.testing.assert_close(forward(model), loss_ref, rtol=1e-4, atol=1e-4) - for site in sites: - mtq.enable_quantizer(model, f"*gdn_{site}_quantizer") + for q in enabled: + q.enable() + for module, policy in zip(gdn_modules, policies): + module.linear_attention_config = policy restored = _make_model(size) sharded_state_dict_test_helper(checkpoint_path, model, restored, forward) restored_gdn = [m for m in restored.modules() if isinstance(m, _QuantGatedDeltaNet)] assert len(restored_gdn) == len(gdn_modules) - for module in restored_gdn: - assert module.gdn_state_quantizer.is_enabled == ("state" in sites) - assert module.gdn_w_quantizer.is_enabled == ("w" in sites) + for module, original in zip(restored_gdn, gdn_modules): + assert module.linear_attention_config == original.linear_attention_config + assert module.gdn_state_quantizer.is_enabled + assert not module.gdn_w_quantizer.is_enabled + assert module.gdn_state_quantizer.num_bits == original.gdn_state_quantizer.num_bits + assert module.gdn_state_qdq_block_v == 32 assert module.in_proj.weight.grad is not None assert torch.isfinite(module.in_proj.weight.grad).all() @@ -143,16 +190,26 @@ def _test_gdn_qat_helper(rank, size, checkpoint_path, sites): optimizer.step() assert not torch.equal(restored_gdn[0].in_proj.weight, before) + fused_model = _make_model(size) + for module in fused_model.modules(): + if isinstance(module, GatedDeltaNet): + module.gdn_pre_gated_delta_rule_fusion = True + with pytest.raises(NotImplementedError, match="unfused input-preparation hook"): + mtq.quantize(fused_model, _gdn_config()) + -def _compile_gdn_qat_kernels(rank, size, sites): +def _compile_gdn_qat_kernels(rank, size): initialize_for_megatron( tensor_model_parallel_size=size, pipeline_model_parallel_size=1, seed=SEED ) model = _make_model(size) - forward = get_forward(model) + forward = _gdn_forward(model) with torch.no_grad(): forward(model) - mtq.quantize(model, _gdn_config(sites)) + cfg = _gdn_config() + # The functional case switches state tiles before running inference. + cfg["linear_attention"][0]["cfg"]["state_block_v"] = 32 + mtq.quantize(model, cfg) with torch.no_grad(): forward(model) model.train() @@ -163,13 +220,10 @@ def _compile_gdn_qat_kernels(rank, size, sites): @pytest.fixture def compiled_gdn_workers(dist_workers_size_1): """Warm one small QAT model in the same worker, outside the test-call budget.""" - # Ampere exercises W QDQ; native FP8 devices also exercise state QDQ. - sites = ("state", "w") if torch.cuda.get_device_capability() >= (8, 9) else ("w",) - dist_workers_size_1.run(_compile_gdn_qat_kernels, sites) - return dist_workers_size_1, sites + dist_workers_size_1.run(_compile_gdn_qat_kernels) + return dist_workers_size_1 def test_gdn_qat_and_sharded_restore(compiled_gdn_workers, tmp_path): """Train through QDQ after a Megatron distributed-checkpoint round trip.""" - workers, sites = compiled_gdn_workers - workers.run(_test_gdn_qat_helper, tmp_path, sites) + compiled_gdn_workers.run(_test_gdn_qat_helper, tmp_path) diff --git a/tests/gpu_megatron/torch/quantization/plugins/test_megatron_kda.py b/tests/gpu_megatron/torch/quantization/plugins/test_megatron_kda.py new file mode 100644 index 00000000000..79c4b89e56c --- /dev/null +++ b/tests/gpu_megatron/torch/quantization/plugins/test_megatron_kda.py @@ -0,0 +1,174 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from contextlib import nullcontext + +import pytest +import torch +from _test_utils.torch.megatron.utils import ( + initialize_for_megatron, + load_distributed_checkpoint, + save_distributed_checkpoint, +) +from megatron.core import parallel_state +from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.transformer import TransformerConfig + +import modelopt.torch.quantization as mtq +from modelopt.recipe import load_recipe +from modelopt.torch.opt.plugins.mcore_dist_checkpointing import ( + restore_sharded_modelopt_state, + save_sharded_modelopt_state, +) +from modelopt.torch.quantization.linear_attention import ( + LinearAttentionConfig, + linear_attention_training_phase, +) + +KimiDeltaAttention = pytest.importorskip("megatron.core.ssm.gated_delta_net.kda").KimiDeltaAttention +pytest.importorskip("vllm") + + +def _layer(): + config = TransformerConfig( + num_layers=1, + hidden_size=64, + num_attention_heads=2, + linear_num_key_heads=2, + linear_num_value_heads=2, + linear_key_head_dim=32, + linear_value_head_dim=32, + linear_conv_kernel_dim=4, + experimental_attention_variant="kda", + is_hybrid_model=True, + normalization="RMSNorm", + activation_func=torch.nn.functional.silu, + params_dtype=torch.float32, + gradient_accumulation_fusion=False, + recompute_granularity="selective", + recompute_modules=["gdn"], + ) + spec = hybrid_stack_spec.submodules.kda_layer.submodules.self_attention + model = ( + KimiDeltaAttention( + config, + submodules=spec.submodules, + layer_number=1, + pg_collection=ProcessGroupCollection( + tp=parallel_state.get_tensor_model_parallel_group(), + cp=parallel_state.get_context_parallel_group(), + ), + ) + .cuda() + .train() + ) + with torch.no_grad(): + model.A_log.fill_(-4) + model.dt_bias.zero_() + return model + + +def _case(cfg): + initialize_for_megatron(tensor_model_parallel_size=1, pipeline_model_parallel_size=1, seed=73) + model = _layer() + hidden = torch.randn(73, 1, 64, device="cuda", dtype=torch.bfloat16, requires_grad=True) + + def forward(layer): + policy = getattr(layer, "linear_attention_config", None) + phase = ( + linear_attention_training_phase(layer, [31]) + if policy is not None and policy.backend == "serving" + else nullcontext() + ) + with phase, torch.autocast("cuda", dtype=torch.bfloat16): + return layer(hidden, attention_mask=None)[0] + + with torch.no_grad(): + baseline = forward(model) + mtq.quantize(model, cfg) + return model, hidden, forward, baseline + + +def _compile_kda(rank, size, cfg): + model, hidden, forward, _ = _case(cfg) + with linear_attention_training_phase(model, [31]): + forward(model).float().square().mean().backward() + torch.cuda.synchronize() + + +def _test_kda(rank, size, cfg, checkpoint_path): + model, hidden, forward, baseline = _case(cfg) + assert model.kda_state_quantizer.is_enabled + assert model.kda_state_quantizer.num_bits == 8 + assert model.linear_attention_config.precision == "vllm" + assert model.kda_state_quantizer.block_sizes == {-1: 32} + kernel = model.gated_delta_rule + with torch.no_grad(): + assert not torch.equal(forward(model), baseline) + assert model.gated_delta_rule is kernel + + # The sharded checkpoint must restore this per-module change, not recipe defaults. + model.linear_attention_config.state_block_v = 32 + policy = model.linear_attention_config + mtq.disable_quantizer(model, "*") + model.linear_attention_config = LinearAttentionConfig() + with torch.no_grad(): + torch.testing.assert_close(forward(model), baseline, rtol=0, atol=0) + model.linear_attention_config = policy + mtq.enable_quantizer(model, "*kda_state_quantizer") + + with torch.no_grad(): + expected = forward(model) + # A change made after conversion must survive beyond the saved recipe defaults. + mtq.disable_quantizer(model, "*kda_state_quantizer") + save_distributed_checkpoint(checkpoint_path, model) + save_sharded_modelopt_state([model], checkpoint_path) + restored = _layer() + restore_sharded_modelopt_state([restored], checkpoint_path) + load_distributed_checkpoint(checkpoint_path, restored) + assert not restored.kda_state_quantizer.is_enabled + mtq.enable_quantizer(restored, "*kda_state_quantizer") + # Megatron recomputes the core during backward, after forward's kernel wrapper exits. + with linear_attention_training_phase(restored, [31]): + actual = forward(restored) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + actual.float().square().mean().backward() + assert not hasattr(restored, "kda_w_quantizer") + assert restored._linear_attention_prefill_lengths is None + assert restored.linear_attention_config == policy + assert restored.gated_delta_rule is kernel + assert torch.isfinite(hidden.grad).all() + for parameter in restored.parameters(): + assert parameter.grad is not None + assert torch.isfinite(parameter.grad).all() + before = restored.in_proj.weight.detach().clone() + torch.optim.SGD(restored.parameters(), lr=0.1).step() + assert not torch.equal(before, restored.in_proj.weight) + + +@pytest.fixture(scope="module") +def compiled_kda_workers(dist_workers_size_1): + """Warm one KDA shape outside the functional test timer.""" + cfg = load_recipe( + "general/ptq/linear_attention_state_int8_block32_dynamic" + ).quantize.model_dump() + dist_workers_size_1.run(_compile_kda, cfg) + return dist_workers_size_1, cfg + + +def test_kda_qat_and_sharded_restore(compiled_kda_workers, tmp_path): + workers, cfg = compiled_kda_workers + workers.run(_test_kda, cfg, tmp_path) diff --git a/tests/gpu_vllm/torch/quantization/test_linear_attention_replay.py b/tests/gpu_vllm/torch/quantization/test_linear_attention_replay.py new file mode 100644 index 00000000000..d6e696b2b40 --- /dev/null +++ b/tests/gpu_vllm/torch/quantization/test_linear_attention_replay.py @@ -0,0 +1,143 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Minimal native-cache and backward checks for the optional ReplaySSM serving fork.""" + +import pytest +import torch + +from modelopt.torch.quantization.linear_attention import LinearAttentionConfig, recurrent_decode + + +@pytest.fixture( + scope="module", + params=[(False, 1), (False, 4), (True, 4)], + ids=["gdn-token", "gdn-replay", "kda-replay"], +) +def compiled_replay(request): + native = pytest.importorskip( + "vllm.model_executor.layers.fla.ops.fused_recurrent_replayssm", + exc_type=ModuleNotFoundError, + ) + channel, window = request.param + torch.manual_seed(71) + q, k = [ + torch.nn.functional.normalize(torch.randn(5, 1, 64, device="cuda"), dim=-1).bfloat16() + for _ in range(2) + ] + v = torch.randn_like(q) + raw = torch.randn(5, 1, device="cuda", dtype=torch.bfloat16) + raw_beta = torch.randn_like(raw) + rate = torch.full((1,), -3.0, device="cuda") + bias = torch.zeros_like(rate) + g = ( + -torch.rand_like(q, dtype=torch.float32) * 0.03 + if channel + else -rate.exp() * torch.nn.functional.softplus(raw.float()) + ) + beta = raw_beta.float().sigmoid() + args = tuple(x.detach().requires_grad_() for x in (q, k, v, g, beta)) + initial = (torch.randn(1, 64, 64, device="cuda") * 0.1).requires_grad_() + config = LinearAttentionConfig( + backend="serving", + precision="replayssm", + replay_window=window, + ) + kwargs = { + "config": config, + "state_qdq": True, + "state_format": "int8", + "replay_gate_inputs": None if channel else (raw, raw_beta, rate, bias), + } + + def forward(): + return recurrent_decode(*args, initial_state=initial, **kwargs) + + out, carry = forward() + torch.autograd.grad(out.sum() + carry.reconstruct().sum(), (*args, initial)) + torch.cuda.synchronize() + return native, args, initial, kwargs, forward + + +def test_replay_matches_persistent_serving_cache(compiled_replay): + native, args, initial, kwargs, forward = compiled_replay + output, carry = forward() + channel = args[3].ndim == 3 + window = kwargs["config"].replay_window + state = torch.zeros(2, 1, 64, 64, device="cuda", dtype=torch.int8) + scales = torch.zeros(2, 1, 2, 64, device="cuda", dtype=torch.float16) + updates = torch.zeros(2, 1, window, 64, device="cuda", dtype=torch.bfloat16) + keys = torch.zeros_like(updates) + gates = torch.zeros((2, 1, window, 64) if channel else (2, 1, window), device="cuda") + indices = torch.ones(1, device="cuda", dtype=torch.int32) + with torch.no_grad(): + native.prefill_write_checkpoint( + initial.transpose(-1, -2)[None].contiguous(), + state, + None, + None, + scales, + None, + None, + indices, + 8, + hadamard=True, + ) + for t in range(len(args[0])): + q, k, v, gate, beta = (x[t] for x in args) + if channel: + a, b, rate, bias = gate.flatten()[None], beta[None], None, None + else: + raw, raw_beta, rate, bias = kwargs["replay_gate_inputs"] + a, b = raw[t : t + 1], raw_beta[t : t + 1] + expected = torch.empty(1, 1, 64, device="cuda", dtype=torch.bfloat16) + native.fused_recurrent_gated_delta_rule_replayssm( + torch.cat([x.flatten() for x in (q, k, v)])[None], + a, + b, + rate, + bias, + 64**-0.5, + state, + updates, + keys, + gates, + expected, + indices, + torch.tensor([t % window], device="cuda", dtype=torch.int32), + quant_state_bits=8, + state_scale=scales, + hadamard_value_basis=True, + block_v=64, + _is_kda=channel, + ) + torch.testing.assert_close(output[t].bfloat16(), expected[0], rtol=0, atol=0) + decoded = state[1].transpose(-1, -2).float().unflatten(-1, (-1, 32)) + metadata = scales[1].transpose(-1, -2) + torch.testing.assert_close( + carry.anchor.values, (decoded * metadata[..., None]).flatten(-2), rtol=0, atol=0 + ) + torch.testing.assert_close(carry.anchor.scales, metadata, rtol=0, atol=0) + for i, entry in enumerate(carry.entries): + torch.testing.assert_close( + entry.key.values.float(), keys[1, :, i].float(), rtol=0, atol=0 + ) + torch.testing.assert_close( + entry.update.values.float(), updates[1, :, i].float(), rtol=0, atol=0 + ) + grads = torch.autograd.grad( + output.square().mean() + carry.reconstruct().square().mean(), (*args, initial) + ) + assert all(torch.isfinite(grad).all() and grad.abs().sum() > 0 for grad in grads) diff --git a/tests/gpu_vllm/torch/quantization/test_linear_attention_training.py b/tests/gpu_vllm/torch/quantization/test_linear_attention_training.py new file mode 100644 index 00000000000..785a7b530d4 --- /dev/null +++ b/tests/gpu_vllm/torch/quantization/test_linear_attention_training.py @@ -0,0 +1,130 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from functools import partial + +import pytest +import torch + +from modelopt.torch.quantization.config import QuantizerAttributeConfig +from modelopt.torch.quantization.linear_attention import ( + LinearAttentionConfig, + gdn_state_qat, + kda_state_qat, +) +from modelopt.torch.quantization.nn import TensorQuantizer + +vllm = pytest.importorskip("vllm") + +from modelopt.torch.kernels.quantization.linear_attention.serving._compat import fla_module + + +@torch.no_grad() +def _native_forward(args, quantizer, kda): + """Execute native prefill/decode directly, with QDQ in ModelOpt's [K,V] basis.""" + prefill = fla_module("kda").chunk_kda if kda else fla_module("chunk").chunk_gated_delta_rule + step = ( + fla_module("kda").fused_recurrent_kda + if kda + else fla_module("fused_recurrent").fused_recurrent_gated_delta_rule + ) + value_first = vllm.__version_tuple__[:2] >= (0, 16) + + def layout(state): + return state.transpose(-1, -2).contiguous() if value_first else state + + state = layout(args[0].new_zeros(1, 1, args[0].shape[-1], args[2].shape[-1]).float()) + prefix, state = prefill( + *[x[:, :65].detach().clone() for x in args], + initial_state=state, + output_final_state=True, + use_qk_l2norm_in_kernel=True, + cu_seqlens=torch.tensor([0, 65], device="cuda", dtype=torch.int32), + ) + outputs = [prefix] + state = layout(quantizer(layout(state)[0])[None]) + for token in range(65, 73): + output, state = step( + *[x[:, token : token + 1].detach().clone() for x in args], + initial_state=state, + inplace_final_state=False, + use_qk_l2norm_in_kernel=True, + cu_seqlens=torch.tensor([0, 1], device="cuda", dtype=torch.int32), + ) + outputs.append(output) + state = layout(quantizer(layout(state)[0])[None]) + return torch.cat(outputs, dim=1), layout(state) + + +@pytest.fixture(scope="module", params=[False, True]) +def compiled_serving_case(request): + """Compile one shared BF16 shape per model, outside the test-call timer.""" + kda = request.param + torch.manual_seed(73) + # A rectangular GDN state detects accidental key/value-axis swaps. + args = [ + torch.randn(1, 73, 1, dim, device="cuda", dtype=torch.bfloat16) + for dim in (32, 32, 32 if kda else 64) + ] + args += [ + -torch.rand((1, 73, 1, 32) if kda else (1, 73, 1), device="cuda") * 0.03, + torch.rand(1, 73, 1, device="cuda") * 0.4, + ] + args = [x.requires_grad_() for x in args] + policy = LinearAttentionConfig(backend="serving", precision="vllm") + quantizer = TensorQuantizer( + QuantizerAttributeConfig( + num_bits=8, + type="dynamic", + block_sizes={-1: 32}, + narrow_range=True, + pass_through_bwd=True, + ) + ).cuda() + forward = partial( + kda_state_qat if kda else gdn_state_qat, + *args, + policy=policy, + state_quantizer=quantizer, + prefill_lengths=[65], + use_qk_l2norm_in_kernel=True, + output_final_state=True, + ) + output, state = forward() + torch.autograd.grad(output.float().sum() + state.sum(), args) + expected = _native_forward(args, quantizer, kda) + torch.cuda.synchronize() + return args, quantizer, forward, expected + + +def test_serving_state_qdq_and_handoff_gradient(compiled_serving_case): + args, quantizer, forward, (expected, expected_state) = compiled_serving_case + calls = [] + handle = quantizer.register_forward_hook(lambda *_: calls.append(True)) + output, state = forward() + handle.remove() + # One handoff QDQ, then one QDQ per suffix update; none inside the fresh prefill. + assert len(calls) == 9 + with torch.no_grad(): + torch.testing.assert_close(output, expected, rtol=0, atol=0) + torch.testing.assert_close(state, expected_state, rtol=0, atol=0) + quantizer.disable() + plain, _ = forward() + assert not torch.equal(output[:, 65:], plain[:, 65:]) + gradients = torch.autograd.grad( + output[:, 65:].float().square().sum() + state.square().sum(), args + ) + assert all(torch.isfinite(x).all() for x in gradients) + assert gradients[1][:, :65].abs().sum() > 0 diff --git a/tests/gpu_vllm/torch/quantization/test_vllm_linear_attention.py b/tests/gpu_vllm/torch/quantization/test_vllm_linear_attention.py new file mode 100644 index 00000000000..6f4a58e0e0e --- /dev/null +++ b/tests/gpu_vllm/torch/quantization/test_vllm_linear_attention.py @@ -0,0 +1,244 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Native vLLM state-boundary QDQ and policy restoration on tiny offline models.""" + +import copy +import gc +import importlib.util +import math +import os +import shutil +from pathlib import Path + +import pytest +import torch +from packaging.version import Version +from transformers import Qwen3NextConfig +from vllm import LLM, SamplingParams +from vllm import __version__ as vllm_version +from vllm.transformers_utils.configs.kimi_linear import KimiLinearConfig + +import modelopt.torch.opt as mto +from modelopt.torch.opt.conversion import ModeloptStateManager +from modelopt.torch.quantization.conversion import restore_quantizer_state +from modelopt.torch.quantization.linear_attention import LinearAttentionConfig +from modelopt.torch.quantization.plugins.vllm_linear_attention import _QuantVllmLinearAttention + +ROOT = Path(__file__).parents[4] + + +def _example_module(name): + spec = importlib.util.spec_from_file_location( + name + "_test", ROOT / "examples/vllm_serve" / (name + ".py") + ) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def _state_worker(worker, *, action="audit"): + model = worker.model_runner.model + if hasattr(model, "unwrap"): + model = model.unwrap() + layers = [m for m in model.modules() if isinstance(m, _QuantVllmLinearAttention)] + assert layers + if action == "restore": + state = worker._state_checkpoint + manager = ModeloptStateManager() + manager.load_state_dict(state["modelopt_state_dict"], state["modelopt_version"]) + for _, config, metadata in manager.modes_with_states(): + restore_quantizer_state(model, config, metadata) + for layer in layers: + if action == "disable": + layer._linear_attn_state.disable() + layer.linear_attention_config = LinearAttentionConfig() + elif action == "setup": + layer.linear_attention_config.state_block_v = 16 + layer._state_calls = {"prefill": 0, "decode": 0, "changed": 0} + original = layer._quantized_state_call + + def checked_call(native, *args, _layer=layer, _original=original, **kwargs): + state = kwargs["initial_state"] + indices = kwargs.get("ssm_state_indices") + if indices is not None: + indices = indices[: kwargs["cu_seqlens"].numel() - 1].long() + selected = state if indices is None else state.index_select(0, indices) + quantizer = copy.deepcopy(_layer._linear_attn_state) + rounded = torch.cat([quantizer(x) for x in selected.split(16, -1)], -1) + expected = state.clone() + if indices is None: + expected = rounded + else: + expected.index_copy_(0, indices, rounded) + + def checked_native(*a, **kw): + # Check the full cache, including inactive slots, before delegating. + torch.testing.assert_close(kw["initial_state"], expected, atol=0, rtol=0) + return native(*a, **kw) + + _layer._state_calls["prefill" if indices is None else "decode"] += 1 + _layer._state_calls["changed"] += int(not torch.equal(selected, rounded)) + return _original(checked_native, *args, **kwargs) + + layer._quantized_state_call = checked_call + if action == "setup": + worker._state_checkpoint = copy.deepcopy(mto.modelopt_state(model)) + return [dict(layer._state_calls) for layer in layers] + + +def _tiny_model(path, kind): + common = { + "torch_dtype": "bfloat16", + "hidden_size": 128, + "intermediate_size": 256, + "num_attention_heads": 4, + "head_dim": 32, + "vocab_size": 512, + "max_position_embeddings": 128, + } + if kind == "gdn": + config = Qwen3NextConfig( + **common, + num_hidden_layers=2, + num_key_value_heads=2, + linear_num_value_heads=4, + linear_num_key_heads=2, + linear_key_head_dim=32, + linear_value_head_dim=64, + num_experts=4, + num_experts_per_tok=2, + moe_intermediate_size=64, + shared_expert_intermediate_size=64, + mlp_only_layers=[], + full_attention_interval=2, + architectures=["Qwen3NextForCausalLM"], + ) + else: + config = KimiLinearConfig( + **common, + num_hidden_layers=1, + architectures=["KimiLinearForCausalLM"], + linear_attn_config={ + "head_dim": 32, + "num_heads": 4, + "short_conv_kernel_size": 4, + "kda_layers": [1], + "full_attn_layers": [], + }, + ) + config.save_pretrained(path) + shutil.copytree(ROOT / "tests/_test_utils/torch/tokenizer", path, dirs_exist_ok=True) + + +def _generate(llm): + outputs = llm.generate( + [{"prompt_token_ids": [3] * 73}, {"prompt_token_ids": [4]}], + SamplingParams(temperature=0, max_tokens=3, ignore_eos=True, logprobs=1), + use_tqdm=False, + ) + values = [(o.outputs[0].token_ids, o.outputs[0].cumulative_logprob) for o in outputs] + assert all(len(tokens) == 3 and math.isfinite(score) for tokens, score in values) + return values + + +@pytest.fixture(params=["gdn", "kda"]) +def compiled_worker(request, tmp_path, monkeypatch): + """Compile one tiny native model outside the functional test's timeout.""" + if Version(vllm_version).release[:2] != (0, 15): + pytest.skip("The state adapter requires vLLM 0.15.x") + if not torch.cuda.is_available(): + pytest.skip("Requires a CUDA GPU") + monkeypatch.setenv("VLLM_WORKER_MULTIPROC_METHOD", "spawn") + paths = [str(ROOT), str(ROOT / "examples/vllm_serve"), str(Path(__file__).parent)] + for path in reversed(paths): + monkeypatch.syspath_prepend(path) + monkeypatch.setenv("PYTHONPATH", os.pathsep.join([*paths, os.environ.get("PYTHONPATH", "")])) + monkeypatch.setenv( + "RECIPE_PATH", str(ROOT / "examples/vllm_serve/linear_attention_state_int8.yaml") + ) + monkeypatch.delenv("MODELOPT_STATE_PATH", raising=False) + model_path = tmp_path / "model" + _tiny_model(model_path, request.param) + llm = LLM( + model=str(model_path), + load_format="dummy", + dtype="bfloat16", + max_model_len=128, + max_num_seqs=4, + max_num_batched_tokens=64, + enforce_eager=True, + enable_prefix_caching=False, + enable_chunked_prefill=True, + async_scheduling=False, + mamba_cache_dtype="float32", + kv_cache_memory_bytes=128 * 1024**2, + gpu_memory_utilization=0.1, + worker_cls="fakequant_worker.FakeQuantWorker", + seed=17, + ) + try: + llm.collective_rpc(_state_worker, kwargs={"action": "setup"}) + _generate(llm) + yield llm + finally: + llm.llm_engine.engine_core.shutdown() + del llm + gc.collect() + + +def test_state_qdq_and_restore(compiled_worker): + llm = compiled_worker + expected = _generate(llm) + reports = llm.collective_rpc(_state_worker) + assert all(all(count > 0 for count in layer.values()) for rank in reports for layer in rank) + llm.collective_rpc(_state_worker, kwargs={"action": "disable"}) + _generate(llm) + assert llm.collective_rpc(_state_worker) == reports + llm.collective_rpc(_state_worker, kwargs={"action": "restore"}) + assert _generate(llm) == expected + + +def test_saved_linear_attention_policy_and_quantizer_names_follow_mapper(): + module = _example_module("vllm_reload_utils") + + metadata = { + "linear_attention": {"backbone.layers.0.attn": {"backend": "serving", "state_block_v": 16}}, + "quantizer_state": {"backbone.layers.0.attn.kda_state_quantizer": {"_disabled": False}}, + } + state = { + "modelopt_state_dict": [("quantize", {"config": {"quant_cfg": []}, "metadata": metadata})] + } + + def mapper(values): + return { + k.replace("backbone.", "model.").replace(".attn.", ".self_attn."): v + for k, v in values.items() + } + + state["modelopt_state_dict"].append( + ("quantize_algo", {"config": {}, "metadata": copy.deepcopy(metadata)}) + ) + result = module.convert_modelopt_state_to_vllm(state, mapper) + converted = result["modelopt_state_dict"][0][1] + assert list(converted["metadata"]["linear_attention"]) == ["model.layers.0.self_attn"] + assert list(converted["metadata"]["quantizer_state"]) == [ + "model.layers.0.self_attn.kda_state_quantizer" + ] + assert converted["config"]["linear_attention"][0] == { + "module_name": "model.layers.0.self_attn", + "cfg": {"backend": "serving", "state_block_v": 16}, + } + assert "linear_attention" not in result["modelopt_state_dict"][1][1]["config"] diff --git a/tests/unit/torch/quantization/plugins/test_gated_delta_net.py b/tests/unit/torch/quantization/plugins/test_gated_delta_net.py deleted file mode 100644 index 3793edcb189..00000000000 --- a/tests/unit/torch/quantization/plugins/test_gated_delta_net.py +++ /dev/null @@ -1,276 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from copy import deepcopy -from functools import partial - -import pytest -import torch -import torch.nn as nn - -import modelopt.torch.opt as mto -import modelopt.torch.quantization as mtq -from modelopt.torch.quantization.nn import QuantModuleRegistry -from modelopt.torch.quantization.plugins import gated_delta_net -from modelopt.torch.quantization.plugins.gated_delta_net import GatedDeltaNetStateQuantMixin - -GDN_STATE_FP8_DYNAMIC = {"num_bits": (4, 3), "axis": (0, 1), "type": "dynamic"} - - -def chunk_gated_delta_rule(q, k, v, g, beta, **kwargs): - """CPU stand-in for the optional FLA kernel.""" - return q + k + v, None - - -@pytest.fixture(autouse=True) -def mock_fla_kernel(monkeypatch): - monkeypatch.setattr( - gated_delta_net, "_fla_chunk_gated_delta_rule", lambda: chunk_gated_delta_rule - ) - - -class TinyGatedDeltaNet(nn.Module): - """A module that, like Megatron-Core's GatedDeltaNet, calls ``self.gated_delta_rule``.""" - - def __init__(self): - super().__init__() - self.proj = nn.Linear(4, 4) - self.gated_delta_rule = chunk_gated_delta_rule - - def forward(self, x): - out, _ = self.gated_delta_rule(x, x, x, x[..., 0], x[..., 0]) - return self.proj(out) - - -@QuantModuleRegistry.register({TinyGatedDeltaNet: "TinyGatedDeltaNet"}) -class _QuantTinyGatedDeltaNet(GatedDeltaNetStateQuantMixin): - def forward(self, x): - gated_delta_rule = self.gated_delta_rule - self.gated_delta_rule = lambda *a, **kw: self._state_quantized_chunk_gated_delta_rule( - gated_delta_rule, *a, **kw - ) - try: - return super().forward(x) - finally: - self.gated_delta_rule = gated_delta_rule - - -GDN_W_FP8_DYNAMIC = {"num_bits": (4, 3), "axis": (0, 1, 2), "type": "dynamic"} - - -def quant_cfg(state=True, w=False): - entries = [{"quantizer_name": "*", "enable": False}] - if state: - entries.append({"quantizer_name": "*gdn_state_quantizer", "cfg": GDN_STATE_FP8_DYNAMIC}) - if w: - entries.append({"quantizer_name": "*gdn_w_quantizer", "cfg": GDN_W_FP8_DYNAMIC}) - return {"quant_cfg": entries, "algorithm": "max"} - - -@pytest.mark.parametrize( - "attributes", - [ - {"num_bits": (4, 3), "axis": (0, 1)}, # static - {"num_bits": (4, 3), "type": "dynamic"}, # per tensor - {"num_bits": 8, "axis": (0, 1), "type": "dynamic"}, # int8 - {"num_bits": (4, 3), "type": "dynamic", "block_sizes": {-1: 16}}, # blockwise - ], -) -def test_validate_state_quantizer_rejects_unsupported(attributes): - with pytest.raises(ValueError, match="supports only"): - mtq.quantize( - TinyGatedDeltaNet(), - { - "quant_cfg": [ - {"quantizer_name": "*", "enable": False}, - {"quantizer_name": "*gdn_state_quantizer", "cfg": attributes}, - ], - "algorithm": None, - }, - ) - - -def test_dynamic_export_removes_linear_attention_attributes(): - model = _QuantTinyGatedDeltaNet.convert(TinyGatedDeltaNet()) - model.export() - assert type(model) is TinyGatedDeltaNet - for name in ("gdn_state_quantizer", "gdn_w_quantizer"): - assert not hasattr(model, name) - - -def test_disabled_state_quantizer_calls_original_kernel(): - model = TinyGatedDeltaNet() - x = torch.randn(2, 8, 3, 4) - expected = model(x) - - disable_all = {"quant_cfg": [{"quantizer_name": "*", "enable": False}], "algorithm": "max"} - mtq.quantize(model, disable_all, lambda m: m(x)) - - assert isinstance(model, _QuantTinyGatedDeltaNet) - assert not model.gdn_state_quantizer.is_enabled and not model.gdn_w_quantizer.is_enabled - assert torch.equal(model(x), expected) - assert model.gated_delta_rule is chunk_gated_delta_rule, "the kernel swap must be undone" - - -@pytest.mark.parametrize("partial_kernel", [False, True]) -def test_enabled_state_quantizer_uses_state_qdq_kernel(monkeypatch, partial_kernel): - calls = [] - - def fake_state_qdq_kernel(*args, **kwargs): - calls.append(kwargs) - return chunk_gated_delta_rule(*args) - - monkeypatch.setattr( - gated_delta_net, "_state_qdq_chunk_gated_delta_rule", lambda: fake_state_qdq_kernel - ) - model = TinyGatedDeltaNet() - if partial_kernel: - model.gated_delta_rule = partial(chunk_gated_delta_rule, output_final_state=True) - x = torch.randn(2, 8, 3, 4) - mtq.quantize(model, quant_cfg(), lambda m: m(x)) - - model(x) - assert calls and calls[-1] == { - "chunk_size": 64, - "state_qdq": 1, - "state_qdq_block_v": 64, - "w_quantizer": None, - **({"output_final_state": True} if partial_kernel else {}), - } - - # An unrelated callable must not pass validation merely by copying the FLA name. - model.gated_delta_rule = lambda *a, **kw: chunk_gated_delta_rule(*a, **kw) - model.gated_delta_rule.__name__ = "chunk_gated_delta_rule" - with pytest.raises(NotImplementedError, match="supports only FLA"): - model(x) - - -@pytest.mark.parametrize("state", [False, True]) -def test_w_quantizer_is_passed_to_the_kernel(monkeypatch, state): - """``*gdn_w_quantizer`` in the config hands the module's TensorQuantizer to the kernel, with - or without the state quantizer.""" - calls = [] - - def fake_state_qdq_kernel(*args, **kwargs): - calls.append(kwargs) - return chunk_gated_delta_rule(*args) - - monkeypatch.setattr( - gated_delta_net, "_state_qdq_chunk_gated_delta_rule", lambda: fake_state_qdq_kernel - ) - model = TinyGatedDeltaNet() - x = torch.randn(2, 8, 3, 4) - mtq.quantize(model, quant_cfg(state=state, w=True), lambda m: m(x)) - assert model.gdn_w_quantizer.is_enabled and model.gdn_state_quantizer.is_enabled == state - - model(x) - assert calls[-1]["state_qdq"] == int(state) - assert calls[-1]["w_quantizer"] is model.gdn_w_quantizer - - # The w quantizer really quantizes: 256 random values per token collapse onto the E4M3 grid, - # which has at most 127 distinct magnitudes per (row-specific) scale. - w = torch.randn(1, 1, 1, 256) - quantized = model.gdn_w_quantizer(w) - assert not torch.equal(quantized, w) - assert torch.unique(quantized.abs()).numel() <= 127 < torch.unique(w.abs()).numel() - - -@pytest.mark.parametrize("site", ["state", "w"]) -@pytest.mark.parametrize( - "overrides", - [{"pass_through_bwd": False}, {"type": "static"}, {"fake_quant": False}, {"rotate": True}], -) -def test_unsupported_quantizer_fails_during_conversion(site, overrides): - cfg = quant_cfg(state=site == "state", w=site == "w") - cfg = deepcopy(cfg) - cfg["quant_cfg"][-1]["cfg"].update(overrides) - with pytest.raises(ValueError, match="supports only"): - mtq.quantize(TinyGatedDeltaNet(), cfg) - - -@pytest.mark.parametrize("axis", [None, (0, 1), (0, 1, 2)]) -def test_w_grouping_is_preserved_during_conversion(axis): - cfg = deepcopy(quant_cfg(state=False, w=True)) - cfg["algorithm"] = None - cfg["quant_cfg"][-1]["cfg"]["axis"] = axis - model = mtq.quantize(TinyGatedDeltaNet(), cfg) - assert model.gdn_w_quantizer.axis == axis - - -@pytest.mark.parametrize(("state", "w"), [(True, False), (False, True), (True, True)]) -def test_quantizer_roundtrip_and_hybrid_selection(tmp_path, state, w): - model = nn.Sequential(TinyGatedDeltaNet(), nn.Linear(4, 4)) - cfg = quant_cfg(state=state, w=w) - cfg["algorithm"] = None - mtq.quantize(model, cfg) - if w: - # Save the actual quantizer settings, including edits after conversion. - model[0].gdn_w_quantizer.axis = None - path = tmp_path / "gdn.pth" - mto.save(model, path) - restored = nn.Sequential(TinyGatedDeltaNet(), nn.Linear(4, 4)) - mto.restore(restored, path) - for name, enabled in (("gdn_state_quantizer", state), ("gdn_w_quantizer", w)): - original = getattr(model[0], name) - quantizer = getattr(restored[0], name) - assert quantizer.is_enabled == enabled - assert quantizer.axis == original.axis - assert quantizer.num_bits == original.num_bits - assert quantizer._dynamic == original._dynamic - assert not hasattr(restored[1], name) - - -def test_quant_cfg_refinement_updates_and_validates_existing_quantized_module(): - cfg = quant_cfg() - cfg["algorithm"] = None - model = mtq.quantize(TinyGatedDeltaNet(), cfg) - assert model.gdn_state_quantizer.is_enabled - assert not model.gdn_w_quantizer.is_enabled - - cfg = quant_cfg(state=False, w=True) - cfg["algorithm"] = None - mtq.quantize(model, cfg) - assert not model.gdn_state_quantizer.is_enabled - assert model.gdn_w_quantizer.is_enabled - - cfg = deepcopy(cfg) - cfg["quant_cfg"][-1]["cfg"]["pass_through_bwd"] = False - with pytest.raises(ValueError, match="supports only"): - mtq.quantize(model, cfg) - - -def test_restore_legacy_gdn_without_new_quantizer_handles(): - config = {"quant_cfg": [{"quantizer_name": "*", "enable": False}], "algorithm": None} - model = mtq.quantize(TinyGatedDeltaNet(), config) - state = mto.modelopt_state(model) - for _, mode_state in state["modelopt_state_dict"]: - metadata = mode_state["metadata"] - for name in ("gdn_state_quantizer", "gdn_w_quantizer"): - metadata["quantizer_state"].pop(name, None) - restored = TinyGatedDeltaNet() - mto.restore_from_modelopt_state(restored, state) - restored.load_state_dict(model.state_dict()) - x = torch.randn(2, 8, 3, 4) - torch.testing.assert_close(restored(x), model(x)) - assert not restored.gdn_state_quantizer.is_enabled - assert not restored.gdn_w_quantizer.is_enabled - - -def test_standard_projection_recipe_leaves_gdn_emulation_disabled(): - model = mtq.quantize( - TinyGatedDeltaNet(), mtq.FP8_DEFAULT_CFG, lambda m: m(torch.randn(2, 8, 3, 4)) - ) - assert not model.gdn_state_quantizer.is_enabled - assert not model.gdn_w_quantizer.is_enabled diff --git a/tests/unit/torch/quantization/plugins/test_gdn.py b/tests/unit/torch/quantization/plugins/test_gdn.py new file mode 100644 index 00000000000..c701b7c1e7b --- /dev/null +++ b/tests/unit/torch/quantization/plugins/test_gdn.py @@ -0,0 +1,348 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from copy import deepcopy + +import pytest +import torch +import torch.nn as nn +from _test_utils.torch.distributed.utils import spawn_multiprocess_job + +import modelopt.torch.opt as mto +import modelopt.torch.quantization as mtq +from modelopt.torch.quantization.config import QuantizeConfig +from modelopt.torch.quantization.linear_attention import ( + LinearAttentionConfig, + linear_attention_training_phase, +) +from modelopt.torch.quantization.nn import QuantModuleRegistry +from modelopt.torch.quantization.plugins import gdn +from modelopt.torch.quantization.plugins.gdn import GatedDeltaNetStateQuantMixin +from modelopt.torch.quantization.plugins.kda import KimiDeltaAttentionStateQuantMixin + +GDN_STATE_FP8_DYNAMIC = {"num_bits": (4, 3), "axis": (0, 1), "type": "dynamic"} + + +def chunk_gated_delta_rule(q, k, v, g, beta, **kwargs): + """CPU stand-in for the optional FLA kernel.""" + return q + k + v, None + + +@pytest.fixture(autouse=True) +def mock_fla_kernel(monkeypatch): + monkeypatch.setattr(gdn, "_fla_chunk_gated_delta_rule", lambda: chunk_gated_delta_rule) + + +class TinyGatedDeltaNet(nn.Module): + """A module that, like Megatron-Core's GatedDeltaNet, calls ``self.gated_delta_rule``.""" + + def __init__(self): + super().__init__() + self.proj = nn.Linear(4, 4) + self.gated_delta_rule = chunk_gated_delta_rule + + def forward(self, x): + out, _ = self.gated_delta_rule(x, x, x, x[..., 0], x[..., 0]) + return self.proj(out) + + +@QuantModuleRegistry.register({TinyGatedDeltaNet: "TinyGatedDeltaNet"}) +class _QuantTinyGatedDeltaNet(GatedDeltaNetStateQuantMixin): + def forward(self, x): + gated_delta_rule = self.gated_delta_rule + self.gated_delta_rule = lambda *a, **kw: self._state_quantized_chunk_gated_delta_rule( + gated_delta_rule, *a, **kw + ) + try: + return super().forward(x) + finally: + self.gated_delta_rule = gated_delta_rule + + +def quant_cfg(): + return { + "quant_cfg": [ + {"quantizer_name": "*", "enable": False}, + {"quantizer_name": "*gdn_state_quantizer", "cfg": GDN_STATE_FP8_DYNAMIC}, + ], + "linear_attention": [{"module_name": "*", "cfg": {"backend": "serving"}}], + "algorithm": "max", + } + + +@pytest.mark.parametrize( + "attributes", + [ + {"num_bits": (4, 3), "axis": (0, 1)}, # static + {"num_bits": (4, 3), "type": "dynamic"}, # per tensor + {"num_bits": 8, "axis": (0, 1), "type": "dynamic"}, # int8 + {"num_bits": (4, 3), "type": "dynamic", "block_sizes": {-1: 16}}, # blockwise + ], +) +def test_validate_state_quantizer_rejects_unsupported(attributes): + with pytest.raises(ValueError, match="supports only"): + mtq.quantize( + TinyGatedDeltaNet(), + { + "quant_cfg": [ + {"quantizer_name": "*", "enable": False}, + {"quantizer_name": "*gdn_state_quantizer", "cfg": attributes}, + ], + "algorithm": None, + }, + ) + + +def test_dynamic_export_removes_linear_attention_attributes(): + model = _QuantTinyGatedDeltaNet.convert(TinyGatedDeltaNet()) + model.export() + assert type(model) is TinyGatedDeltaNet + for name in ( + "gdn_state_quantizer", + "gdn_w_quantizer", + "linear_attention_config", + "_linear_attention_prefill_lengths", + ): + assert not hasattr(model, name) + + +def test_disabled_state_quantizer_calls_original_kernel(): + model = TinyGatedDeltaNet() + x = torch.randn(2, 8, 3, 4) + expected = model(x) + + disable_all = {"quant_cfg": [{"quantizer_name": "*", "enable": False}], "algorithm": "max"} + mtq.quantize(model, disable_all, lambda m: m(x)) + + assert isinstance(model, _QuantTinyGatedDeltaNet) + assert not model.gdn_state_quantizer.is_enabled and not model.gdn_w_quantizer.is_enabled + assert torch.equal(model(x), expected) + assert model.gated_delta_rule is chunk_gated_delta_rule, "the kernel swap must be undone" + + +@pytest.mark.parametrize("mixin", [GatedDeltaNetStateQuantMixin, KimiDeltaAttentionStateQuantMixin]) +def test_state_qat_requires_explicit_serving_policy(mixin): + model = mixin.convert(TinyGatedDeltaNet()) + model._linear_attn_state.set_from_attribute_config(GDN_STATE_FP8_DYNAMIC) + model._linear_attn_state.enable() + with pytest.raises(ValueError, match="requires backend='serving'"): + model.validate_linear_attention() + + +def test_training_phase_routes_prefill_lengths(monkeypatch): + cfg = {**quant_cfg(), "algorithm": None} + missing_policy = {key: value for key, value in cfg.items() if key != "linear_attention"} + with pytest.raises(ValueError, match="requires backend='serving'"): + mtq.quantize(TinyGatedDeltaNet(), missing_policy) + model = mtq.quantize(TinyGatedDeltaNet(), cfg) + x = torch.randn(2, 8, 3, 4) + calls = [] + + def forward(*args, **kwargs): + calls.append(kwargs) + return chunk_gated_delta_rule(*args) + + monkeypatch.setattr(gdn, "gdn_state_qat", forward) + with linear_attention_training_phase(model, [4, 4]): + model(x) + assert calls[-1]["prefill_lengths"] == (4, 4) + + +@pytest.mark.parametrize("phase", ["convert", "forward", "restore"]) +def test_enabled_legacy_w_quantizer_is_rejected(phase): + cfg = {"quant_cfg": [{"quantizer_name": "*", "enable": False}], "algorithm": None} + model = mtq.quantize(TinyGatedDeltaNet(), cfg) + with pytest.raises(ValueError, match="GDN W quantization is no longer supported"): + if phase == "convert": + cfg["quant_cfg"].append({"quantizer_name": "*gdn_w_quantizer", "enable": True}) + mtq.quantize(TinyGatedDeltaNet(), cfg) + else: + model.gdn_w_quantizer.enable() + if phase == "forward": + model(torch.randn(2, 8, 3, 4)) + else: + mto.restore_from_modelopt_state(TinyGatedDeltaNet(), mto.modelopt_state(model)) + + +@pytest.mark.parametrize( + "overrides", + [{"pass_through_bwd": False}, {"type": "static"}, {"fake_quant": False}, {"rotate": True}], +) +def test_unsupported_quantizer_fails_during_conversion(overrides): + cfg = deepcopy(quant_cfg()) + cfg["quant_cfg"][-1]["cfg"].update(overrides) + with pytest.raises(ValueError, match="supports only"): + mtq.quantize(TinyGatedDeltaNet(), cfg) + + +def test_quantizer_roundtrip_and_hybrid_selection(tmp_path): + model = nn.Sequential(TinyGatedDeltaNet(), nn.Linear(4, 4)) + cfg = deepcopy(quant_cfg()) + cfg["algorithm"] = None + cfg["quant_cfg"][-1]["cfg"].update( + num_bits=8, axis=None, block_sizes={-1: 32}, unsigned=False, narrow_range=True + ) + cfg["linear_attention"] = [ + {"module_name": "*", "cfg": {"backend": "serving", "state_block_v": 128}}, + {"module_name": "0", "cfg": {"backend": "serving", "state_block_v": 32}}, + ] + mtq.quantize(model, cfg) + assert model[0].gdn_state_qdq_block_v == 32 + model[0].linear_attention_config.state_block_v = 16 + sample = torch.randn(2, 4, 19) + expected = model[0].gdn_state_quantizer(sample) + path = tmp_path / "gdn.pth" + mto.save(model, path) + restored = nn.Sequential(TinyGatedDeltaNet(), nn.Linear(4, 4)) + mto.restore(restored, path) + assert restored[0].linear_attention_config == model[0].linear_attention_config + assert restored[0].gdn_state_qdq_block_v == 16 + for name, enabled in (("gdn_state_quantizer", True), ("gdn_w_quantizer", False)): + original = getattr(model[0], name) + quantizer = getattr(restored[0], name) + assert quantizer.is_enabled == enabled + assert quantizer.axis == original.axis + assert quantizer.num_bits == original.num_bits + assert quantizer.block_sizes == original.block_sizes + assert quantizer._dynamic == original._dynamic + assert not hasattr(restored[1], name) + torch.testing.assert_close(restored[0].gdn_state_quantizer(sample), expected) + + +@pytest.mark.parametrize("reverse", [False, True], ids=["fp8-to-replay", "replay-to-fp8"]) +def test_restore_changed_policy_and_quantizer_together(tmp_path, reverse): + states = [ + ( + GDN_STATE_FP8_DYNAMIC, + LinearAttentionConfig(backend="serving", precision="vllm_0_15"), + ), + ( + {"num_bits": 8, "axis": (0, 1), "type": "dynamic", "narrow_range": True}, + LinearAttentionConfig(backend="serving", precision="replayssm"), + ), + ] + if reverse: + states.reverse() + (original_quantizer, original_policy), (saved_quantizer, saved_policy) = states + cfg = quant_cfg() + cfg["algorithm"] = None + cfg["quant_cfg"][-1]["cfg"] = original_quantizer + cfg["linear_attention"][0]["cfg"] = original_policy.model_dump() + model = mtq.quantize(TinyGatedDeltaNet(), cfg) + model.gdn_state_quantizer.set_from_attribute_config(saved_quantizer) + model.linear_attention_config = saved_policy + model.validate_linear_attention() + + path = tmp_path / "changed-policy.pth" + mto.save(model, path) + restored = mto.restore(TinyGatedDeltaNet(), path) + assert restored.linear_attention_config == saved_policy + assert restored.gdn_state_quantizer.num_bits == saved_quantizer["num_bits"] + assert restored.gdn_state_quantizer.is_enabled + restored.validate_linear_attention() + + +def test_quant_cfg_refinement_updates_and_validates_existing_quantized_module(): + cfg = quant_cfg() + cfg["algorithm"] = None + model = mtq.quantize(TinyGatedDeltaNet(), cfg) + assert model.gdn_state_quantizer.is_enabled + assert not model.gdn_w_quantizer.is_enabled + + cfg = deepcopy(quant_cfg()) + cfg["algorithm"] = None + cfg["linear_attention"] = [ + {"module_name": "", "cfg": {"backend": "serving", "state_block_v": 32}} + ] + mtq.quantize(model, cfg) + assert model.gdn_state_qdq_block_v == 32 + cfg["linear_attention"].append({"module_name": "", "cfg": {"backend": "serving"}}) + mtq.quantize(model, cfg) + assert model.gdn_state_qdq_block_v == 64 + assert model.gdn_state_quantizer.is_enabled + assert not model.gdn_w_quantizer.is_enabled + + cfg = deepcopy(cfg) + cfg["quant_cfg"][-1]["cfg"]["pass_through_bwd"] = False + with pytest.raises(ValueError, match="supports only"): + mtq.quantize(model, cfg) + + +def test_restore_legacy_gdn_without_new_quantizer_handles(): + config = {"quant_cfg": [{"quantizer_name": "*", "enable": False}], "algorithm": None} + model = mtq.quantize(TinyGatedDeltaNet(), config) + state = mto.modelopt_state(model) + for _, mode_state in state["modelopt_state_dict"]: + mode_state["config"].pop("linear_attention", None) + metadata = mode_state["metadata"] + metadata.pop("linear_attention", None) + for name in ("gdn_state_quantizer", "gdn_w_quantizer"): + metadata["quantizer_state"].pop(name, None) + restored = TinyGatedDeltaNet() + mto.restore_from_modelopt_state(restored, state) + restored.load_state_dict(model.state_dict()) + x = torch.randn(2, 8, 3, 4) + torch.testing.assert_close(restored(x), model(x)) + assert not restored.gdn_state_quantizer.is_enabled + assert not restored.gdn_w_quantizer.is_enabled + + +def test_standard_projection_recipe_leaves_gdn_emulation_disabled(): + model = mtq.quantize( + TinyGatedDeltaNet(), mtq.FP8_DEFAULT_CFG, lambda m: m(torch.randn(2, 8, 3, 4)) + ) + assert not model.gdn_state_quantizer.is_enabled + assert not model.gdn_w_quantizer.is_enabled + + +def test_policy_rejects_unmatched_and_unimplemented_modes(): + with pytest.raises(ValueError, match="matches no supported"): + mtq.quantize( + nn.Linear(4, 4), {"linear_attention": [{"module_name": "*"}], "algorithm": None} + ) + for policy in ( + {"chunk_size": 32}, + {"solve": {"method": "neumann"}}, + {"state": {"mode": "token"}}, + ): + with pytest.raises(ValueError): + QuantizeConfig(linear_attention=[{"module_name": "*", "cfg": policy}]) + + +def _test_policy_matches_across_stages(rank, size): + config = { + "quant_cfg": [{"quantizer_name": "*", "enable": False}], + "linear_attention": [{"module_name": "layer", "cfg": {"backend": "serving"}}], + "algorithm": None, + } + try: + # Only rank 1 owns a selected layer; both stages use the same recipe and phase. + layer = TinyGatedDeltaNet() if rank else nn.Linear(4, 4) + model = mtq.quantize(nn.ModuleDict({"layer": layer}), config) + with linear_attention_training_phase(model, [3]): + if rank: + assert model["layer"]._linear_attention_prefill_lengths == (3,) + if rank: + assert model["layer"]._linear_attention_prefill_lengths is None + + config["linear_attention"][0]["module_name"] = "missing" + with pytest.raises(ValueError, match="matches no supported"): + mtq.quantize(model, config) + finally: + torch.distributed.destroy_process_group() + + +def test_policy_matches_across_stages(skip_on_windows): + spawn_multiprocess_job(2, _test_policy_matches_across_stages, backend="gloo") diff --git a/tests/unit/torch/quantization/test_linear_attention_decode.py b/tests/unit/torch/quantization/test_linear_attention_decode.py new file mode 100644 index 00000000000..3c8058df60c --- /dev/null +++ b/tests/unit/torch/quantization/test_linear_attention_decode.py @@ -0,0 +1,57 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import pytest +import torch + +from modelopt.torch.quantization.config import QuantizerAttributeConfig +from modelopt.torch.quantization.linear_attention import LinearAttentionConfig +from modelopt.torch.quantization.linear_attention.decode import _encode +from modelopt.torch.quantization.nn import TensorQuantizer + + +@pytest.mark.parametrize("state_format", ["fp8_e4m3", "int8"]) +def test_state_qdq_matches_tensor_quantizer(state_format): + torch.manual_seed(762) + # A partial value tile checks grouping and padding against TensorQuantizer. + value = torch.randn(2, 3, 5, 19, requires_grad=True) + cfg = {"num_bits": (4, 3), "type": "dynamic", "axis": (0, 1)} + if state_format == "int8": + cfg.update(num_bits=8, unsigned=False, narrow_range=True) + quantizer = TensorQuantizer(QuantizerAttributeConfig(**cfg)) + expected = torch.cat( + [quantizer(tile.flatten(-2)).reshape_as(tile) for tile in value.split(16, -1)], -1 + ) + encoded = _encode(value, True, 16, state_format=state_format, state_quantizer=quantizer) + torch.testing.assert_close(encoded.values, expected, rtol=0, atol=0) + probe = torch.randn_like(value) + (gradient,) = torch.autograd.grad((encoded.values * probe).sum(), value) + torch.testing.assert_close(gradient, probe, rtol=0, atol=0) + assert encoded.scales.shape == (2, 3, 2) + assert not encoded.scales.requires_grad + + +@pytest.mark.parametrize( + "settings", + [ + {"replay_window": 4}, + {"precision": "replayssm", "replay_window": 0}, + {"precision": "replayssm", "replay_window": 65}, + {"precision": "replayssm", "state_block_v": 16}, + ], +) +def test_unified_policy_rejects_incompatible_settings(settings): + with pytest.raises(ValueError): + LinearAttentionConfig(backend="serving", **settings) diff --git a/tests/unit/torch/quantization/test_linear_attention_hadamard.py b/tests/unit/torch/quantization/test_linear_attention_hadamard.py new file mode 100644 index 00000000000..41acfc74aa3 --- /dev/null +++ b/tests/unit/torch/quantization/test_linear_attention_hadamard.py @@ -0,0 +1,35 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import torch + +from modelopt.torch.quantization.linear_attention.decode import _hadamard32 + + +def _rotate(value): + # Independent dense Sylvester matrix; production uses a butterfly transform. + matrix = value.new_tensor([[(-1) ** (i & j).bit_count() for j in range(32)] for i in range(32)]) + return (value.reshape(*value.shape[:-1], -1, 32) @ (matrix / 32**0.5)).reshape_as(value) + + +def test_hadamard_matches_dense_oracle_and_backward(): + torch.manual_seed(321) + state = torch.randn(2, 4, 64, dtype=torch.float64, requires_grad=True) + actual = _hadamard32(state) + expected = _rotate(state) + torch.testing.assert_close(actual, expected, rtol=1e-10, atol=1e-10) + probe = torch.randn_like(state) + (gradient,) = torch.autograd.grad((actual * probe).sum(), state) + torch.testing.assert_close(gradient, _rotate(probe), rtol=1e-10, atol=1e-10) diff --git a/tests/unit/torch/quantization/test_linear_attention_reference.py b/tests/unit/torch/quantization/test_linear_attention_reference.py index d71818d0163..1744fc07d9c 100644 --- a/tests/unit/torch/quantization/test_linear_attention_reference.py +++ b/tests/unit/torch/quantization/test_linear_attention_reference.py @@ -17,9 +17,9 @@ import torch import torch.nn.functional as F from _test_utils.torch.quantization.linear_attention_reference import ( - chunk_gdn_reference, + chunk_delta_rule_reference, recurrent_delta_rule_reference, - state_fp8_qdq_reference, + state_qdq_reference, ) @@ -34,10 +34,11 @@ def inputs(length=7, kda=False): return [x.requires_grad_() for x in (q, k, v, g, beta)] -@pytest.mark.parametrize("packed", [False, True]) -@pytest.mark.parametrize("state_v_first", [False, True]) -def test_chunk_recurrence_output_state_and_gradients(packed, state_v_first): - args = inputs() +@pytest.mark.parametrize( + ("kda", "packed", "state_v_first"), [(False, True, True), (True, False, False)] +) +def test_chunk_recurrence_output_state_and_gradients(kda, packed, state_v_first): + args = inputs(kda=kda) state = torch.randn(2 if packed else 1, 2, 3, 4, dtype=torch.float64) if state_v_first: state = state.transpose(-1, -2).contiguous() @@ -47,7 +48,7 @@ def test_chunk_recurrence_output_state_and_gradients(packed, state_v_first): "state_v_first": state_v_first, "cu_seqlens": torch.tensor([0, 2, 7]) if packed else None, } - actual = chunk_gdn_reference(*args, chunk_size=3, **kwargs) + actual = chunk_delta_rule_reference(*args, chunk_size=3, **kwargs) expected = recurrent_delta_rule_reference(*args, **kwargs) for a, e in zip(actual, expected): torch.testing.assert_close(a, e, rtol=1e-10, atol=1e-10) @@ -98,11 +99,11 @@ def test_state_qdq_granularity_zero_tail_and_identity_ste(): with torch.no_grad(): state[0].zero_() state[1, ..., :16] *= 10 - quantized = state_fp8_qdq_reference(state, block_v=16) + quantized = state_qdq_reference(state, block_v=16) assert torch.isfinite(quantized).all() assert torch.equal(quantized[0], state[0]) assert not torch.equal(quantized, state) - assert not torch.equal(quantized, state_fp8_qdq_reference(state, block_v=64)) + assert not torch.equal(quantized, state_qdq_reference(state, block_v=64)) probe = torch.randn_like(state) (grad,) = torch.autograd.grad((quantized * probe).sum(), state) torch.testing.assert_close(grad, probe, rtol=0, atol=0) diff --git a/tests/unit/torch/quantization/test_tensor_quantizer_cpu.py b/tests/unit/torch/quantization/test_tensor_quantizer_cpu.py index d1512547643..bf8238cb85c 100644 --- a/tests/unit/torch/quantization/test_tensor_quantizer_cpu.py +++ b/tests/unit/torch/quantization/test_tensor_quantizer_cpu.py @@ -119,3 +119,16 @@ def test_grouped_quantizer_preserves_nested_sequential_state_dict_layout(): ) assert list(grouped.state_dict()) == ["0.0._amax", "0.1._amax", "1._amax"] + + +def test_dynamic_block_quantization_accepts_changing_shapes(): + cfg = QuantizerAttributeConfig( + num_bits=8, type="dynamic", block_sizes={-1: 32}, pass_through_bwd=True + ) + quantizer = TensorQuantizer(cfg) + for shape in ((2, 3, 17), (1, 2, 4, 64), (3, 2, 33)): + value = torch.randn(shape, requires_grad=True) + output = quantizer(value) + torch.testing.assert_close(output, TensorQuantizer(cfg)(value), rtol=0, atol=0) + output.sum().backward() + torch.testing.assert_close(value.grad, torch.ones_like(value), rtol=0, atol=0)