Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion .github/workflows/gpu_tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,8 @@ jobs:
matrix:
include:
- example: gpu
timeout: 60
# Includes dependency builds and cold compilation of the FLA forward/backward tests.
timeout: 75
# Pinned to 26.05: benchmark.py uses trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH,
# which newer TensorRT (26.06) removed. Bump once the source is updated for TensorRT 10.
container_image: nvcr.io/nvidia/pytorch:26.05-py3
Expand Down
2 changes: 2 additions & 0 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,8 @@ repos:
modelopt/onnx/quantization/ort_patching.py|
modelopt/torch/_deploy/utils/onnx_utils.py|
modelopt/torch/export/transformer_engine.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/puzzletron/anymodel/models/gpt_oss/gpt_oss_pruned_to_mxfp4.py|
modelopt/torch/quantization/export_onnx.py|
modelopt/torch/quantization/plugins/attention.py|
Expand Down
3 changes: 3 additions & 0 deletions CHANGELOG.rst
Original file line number Diff line number Diff line change
Expand Up @@ -13,10 +13,13 @@ Changelog

*Quantization*

- Add experimental GDN/KDA decode-aware QAT with FP8 or INT8 recurrent states, optional INT8 value-axis Hadamard encoding, KDA decay rounding, and encoded-update replay. Supply explicit prefix lengths through the training phase context; prefix solves remain exact.

- Add ``layerwise.export_dir``: layerwise calibration writes each decoder layer to its own quantized checkpoint shard as it finishes, so no separate ``export_hf_checkpoint()`` pass is needed and, with ``layerwise.checkpoint_dir``, an interrupted run resumes without redoing finished layers. Calibration writes the layer shards; ``finalize()`` on the exporter left on the model adds the tail shard, the index and the config artifacts, and the checkpoint does not load until it runs. ``examples/hf_ptq`` does this for you. Supports FP8 and NVFP4 on single-process models, resident or offloaded, including multimodal models and models with MTP layers; other formats and placements raise ``NotImplementedError`` before calibration starts.
- 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 <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/onnx_ptq/bevformer>`_ 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.

*Megatron Framework (M-LM / M-Bridge)*

Expand Down
1 change: 1 addition & 0 deletions LICENSE
Original file line number Diff line number Diff line change
Expand Up @@ -251,6 +251,7 @@ the following copyright holders, licensed under the MIT License:
Copyright (c) 2025 sgl-project
Copyright (c) 2026 The DeepSpec Authors
Copyright (c) 2023-2026 The ggml authors
Copyright (c) 2023-2026 Songlin Yang, Yu Zhang, Zhiyuan Li

Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
Expand Down
2 changes: 1 addition & 1 deletion docs/source/_templates/autosummary/module.rst
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
:recursive:
{% for item in modules %}
{% set full_item = fullname + '.' + item.split('.')[-1] %}
{% if ('.plugins.' not in full_item or full_item == 'modelopt.torch.opt.plugins.huggingface') and full_item != 'modelopt.torch.quantization.backends.fp8_per_tensor_gemm' %}
{% if ('.plugins.' not in full_item or full_item == 'modelopt.torch.opt.plugins.huggingface') and full_item != 'modelopt.torch.quantization.backends.fp8_per_tensor_gemm' and not full_item.startswith('modelopt.torch.kernels.quantization.linear_attention.fla_') %}
{{ full_item }}
{% endif %}
{%- endfor %}
Expand Down
1 change: 1 addition & 0 deletions examples/llm_qat/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ For background on how QAT enables low-precision accuracy recovery, see the [QAT/
| Background | How QAT/QAD work and when to use each | \[[Link](#background)\] | \[[docs](https://nvidia.github.io/Model-Optimizer/guides/1_quantization.html)\] |
| 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 | FSDP2 config, YAML options | \[[Link](#advanced-topics)\] | |
| Results | Accuracy benchmarks | \[[Link](#results)\] | |
| Resources | Extra links and references | \[[Link](#resources)\] | |
Expand Down
307 changes: 307 additions & 0 deletions examples/llm_qat/linear_attention/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,307 @@
# Quantization-Aware Training for Linear Attention

This example fine-tunes KDA attention parameters with recurrent-state fake
quantization, then saves a Hugging Face checkpoint with ModelOpt state. GDN/KDA
runtime support includes token writes, ReplaySSM, KDA decay approximation, and
FP8 or INT8 state QDQ. The INT8 recipes enable Hadamard rotation by default.

## Run the example

Install ModelOpt using the [QAT setup instructions](../README.md#quick-start),
then install the example dependencies. Run commands from the repository root.

```bash
pip install -r examples/llm_qat/linear_attention/requirements.txt
python examples/llm_qat/linear_attention/train.py \
--model /path/to/local-kda-model \
--train-data /path/to/train.parquet \
--output /path/to/qat-checkpoint \
--train-steps 1 --length 128 --prefill-tokens 64
```

Use a local model/tokenizer snapshot containing FLA `KimiDeltaAttention` layers,
a Parquet file with a `text` column, and a CUDA GPU. The example requires
`fla-core==0.5.1` and `flash-linear-attention==0.5.1`. If the local model requires
custom Python code, review it and explicitly pass `--trust-remote-code`.

The default [INT8 configuration](configs/kda_decode_state_int8.json) leaves the
64-token prefix state unquantized and applies INT8 QDQ in a 32-value Hadamard
basis during the decode suffix. Value dimensions must be divisible by 32.
Loss uses suffix labels. Only KDA attention parameters are trained, in FP32 under
BF16 autocast; other parameters are frozen in BF16. The output contains model
weights, tokenizer files, and ModelOpt quantizer/policy state. Call
`mto.enable_huggingface_checkpointing()` before reloading it with
`AutoModelForCausalLM.from_pretrained` to restore those quantizers and policies.
It is a floating fake-quantized training checkpoint, not a compressed serving model.
This short example does not measure model-quality recovery or performance.

## Enable state quantization

### What the execution policy controls

An **execution policy** is the layer's `LinearAttentionConfig`: the settings that
determine how the recurrence runs and when it applies quantization. Training
uses three separate inputs:

| Input | Purpose | Supplied through |
| --- | --- | --- |
| `TensorQuantizer` settings | Enable a quantization site and choose its numerical format, scales, and gradient behavior | Recipe `quant_cfg`, such as dynamic INT8 with identity STE |
| Execution policy | Choose token or replay mode, Hadamard rotation, handoff quantization, and the execution backend | Recipe `linear_attention` entries |
| Batch metadata | Specify where each sequence switches from prefill to decode | `linear_attention_training_phase(model, prefill_lengths)` |

The same INT8 quantizer settings can round the state after every token or round
a replay anchor every eight tokens. Those schedules produce different recurrent
states, so the policy must specify which computation training should emulate.

`LinearAttentionConfig` holds the overall backend, chunk size, and state settings.
Its optional `decode` field contains a `LinearAttentionDecodeConfig` for token or
replay mode, Torch or Triton implementation, state codec, decay approximation,
and replay settings. This is one nested configuration: supplying a `decode`
dictionary in the recipe constructs the nested config automatically.

Each `linear_attention` entry uses `module_name` to select attention layers and
`cfg` to specify their policy. During `mtq.quantize`, ModelOpt stores that policy
as each matched layer's `linear_attention_config`. Quantizer settings and the
policy persist through ModelOpt save/restore; batch-specific prefill lengths
must be supplied again for each workload.

### Load a state recipe

State quantizers start disabled. The state recipe enables `*kda_state_quantizer`
with signed narrow-range INT8, dynamic scales, and `axis=(0, 1)`. Use
`*gdn_state_quantizer` for GDN. For E4M3, set the quantizer config to
`{"num_bits": [4, 3], "type": "dynamic", "axis": [0, 1]}` and set
`decode.state_codec="tile"`.
Setting `prefill_state_qdq=True` alone does not enable a quantizer.

Start with the token-state recipe and select one of the schedules below **before**
calling `mtq.quantize`:

```python
import json
from pathlib import Path

recipe = json.loads(
Path("examples/llm_qat/linear_attention/configs/kda_decode_state_int8.json").read_text()
)
policy = recipe["linear_attention"][0]["cfg"]
decode = policy["decode"]
```

For GDN, the complete recipe composes the INT8 quantizer unit with the Hadamard
execution policy:

```python
import modelopt.torch.quantization as mtq
from modelopt.recipe import load_recipe

gdn_recipe = load_recipe("general/ptq/gdn_state_int8_dynamic").quantize
mtq.quantize(gdn_model, gdn_recipe)
```

Use the phase context shown below for each forward/backward. Importing only
`configs/ptq/units/gdn_state_int8_dynamic` configures the quantizer without
selecting Hadamard. The complete GDN recipe uses the Torch implementation;
set `gdn_recipe.linear_attention[0].cfg.decode.implementation="triton"` before
conversion to select the fused decode implementation.

### Why training needs prefill/decode boundaries

One training forward can simulate prompt prefill followed by recurrent decode.
The INT8 + Hadamard recipes apply different state quantization schedules to those
phases, so the workload must specify where each sequence switches to decode.
The training batch supplies all tokens; this simulates decode arithmetic without
running a text-generation loop. Total sequence length alone does not identify
the prompt portion.

For a 128-token sequence with `prefill_lengths=[96]`, the default token-state
recipe runs these steps:

1. **Prefill:** process tokens 0–95 with `prefill_state_qdq=False`, leaving state
unquantized during the prefix.
2. **Handoff:** carry the resulting state into decode, rotate its value dimension
into the Hadamard basis, and apply INT8 QDQ because `quantize_initial=True`.
3. **Decode:** process tokens 96–127 recurrently, applying INT8 QDQ after every
state update. Later tokens consume the rounded state; outputs are transformed
back to the original value basis.

The boundary is independent of `chunk_size=64`: this 96-token prefix contains a
full 64-token chunk and a partial 32-token chunk. The phase switches after token
95, not after each chunk.

Two 128-token sequences can use `prefill_lengths=[64, 96]` in the same batch:
their decode suffixes then contain 64 and 32 tokens, respectively. The next batch
can have different lengths while reusing the same quantization recipe. This is
why the execution policy is saved in `linear_attention_config`, while the lengths
are supplied per batch through `linear_attention_training_phase`. They are not
saved as part of the model's quantization policy. The context selects numerical
phases; the caller still supplies training labels and any loss masking.

The combined prefill/decode path requires explicit lengths, including `[0]` for
decode only or `[T]` for an all-prefix sequence of length `T`. The GDN chunk-only
FLA path described below applies its chunk schedule throughout and needs no
phase context. See the tables below for the corresponding quantization settings.

### Choose the prefill and decode boundaries

The table assumes the state quantizer is enabled. The INT8 recipes default to
`"int8_hadamard32"`; prefix state QDQ requires explicitly selecting `"tile"`.
Prefix lengths are supplied separately through `linear_attention_training_phase`.

| Desired state QDQ | `decode.state_codec` | `decode.mode` | `decode.prefill_state_qdq` | Where rounding occurs |
| --- | --- | --- | --- | --- |
| Token decode only (default) | `"int8_hadamard32"` | `"token"` | `False` | At the first nonempty decode handoff and after every suffix token. |
| Prefill and token decode | `"tile"` | `"token"` | `True` | At prefix initialization, each prefix chunk write, decode handoff, and every suffix token. |
| Replay anchors only | `"int8_hadamard32"` | `"replay"` | `False` | At decode handoff and each replay-window refresh. |
| Prefill and replay anchors | `"tile"` | `"replay"` | `True` | At prefix initialization and chunk writes, then decode handoff and replay-window refreshes. |

To enable **prefill and token decode**:

```python
decode.update(mode="token", replay=None, state_codec="tile", prefill_state_qdq=True)
```

For **token decode only**, keep the supplied INT8 recipe's defaults:
`state_codec="int8_hadamard32"` and `prefill_state_qdq=False`.
Prefix state remains unquantized until it enters the decode path.

To enable **ReplaySSM anchor quantization** with an eight-token window:

```python
decode.update(
mode="replay",
state_codec="int8_hadamard32",
prefill_state_qdq=False,
replay={"window": 8, "factor_qdq": False, "encoding": "once"},
)
```

Set `state_codec="tile"` and `prefill_state_qdq=True` to add prefix state QDQ
to this replay configuration, using unrotated INT8 for both phases.
`factor_qdq=False` above isolates state/anchor quantization. Set it to `True` to
also quantize buffered keys and updates to FP8.

`decode.quantize_initial=True` is the default: it quantizes the incoming state
once at the first nonempty decode handoff. Set it to `False` to skip that initial
rounding while keeping later token writes or anchor refreshes quantized. This
setting does not disable prefix state QDQ. With both phases enabled, prefix-final
rounding and decode-handoff rounding are separate configured events.

`chunk_size=64` counts **tokens per prefill chunk**. `state.block_v=64` counts
**value channels per execution tile**. The tile codec shares a scale across all
key channels and this value tile; Hadamard uses one scale per key channel and
32 values. A replay `window=8` counts **suffix tokens between anchor refreshes**.

### Run only the desired phase

For a batch containing one sequence of `T` tokens, choose the context lengths as
follows; for larger batches, provide one length per sequence:

| Workload | Context argument | Required setting |
| --- | --- | --- |
| Prefill followed by decode | `[64]`, with `T > 64` | Select either prefix setting above. |
| Decode only | `[0]` | State quantizer enabled; token or replay policy. |
| Prefill only | `[T]` | `state_codec="tile"`, `prefill_state_qdq=True`; the decode suffix is empty. |

An empty suffix creates no decode quantization event. The combined interface has
one state quantizer per layer: it does not offer a switch to quantize the prefix
while leaving a **nonempty** decode suffix unquantized. Disable the state quantizer
to disable both state and anchor QDQ; replay factor QDQ has its own toggle.

The `int8_hadamard32` codec applies to decode token writes or replay anchors only.
It requires `prefill_state_qdq=False`; use the tile codec to quantize prefix state.

Save a modified recipe as JSON and pass it to `train.py --quant-config`. `--prefill-tokens` sets the phase split; it does not enable state
quantization. The integration loop below applies the configured `recipe` directly.

### GDN chunk-only FLA training

For GDN's existing chunked FLA path, enable the GDN state quantizer and omit the
`linear_attention` execution-policy entry:

```python
gdn_recipe = {
"quant_cfg": [
{"quantizer_name": "*", "enable": False},
{
"quantizer_name": "*gdn_state_quantizer",
"cfg": recipe["quant_cfg"][1]["cfg"],
},
],
"algorithm": None,
}
# Apply mtq.quantize(gdn_model, gdn_recipe), then use normal forward/backward.
```

This explicitly selects unrotated tile QDQ. It rounds the initial state and each
64-token chunk's final state, including a partial final chunk. It needs no phase
context and leaves W/projection quantizers
disabled. KDA uses the materialized backend and can use the all-prefix context
shown above.

## Integrate with a training loop

Apply the configured `recipe` above with `mtq.quantize`, then supply one prefix length per sequence.
Keep the phase context active through backward so activation-checkpoint
recomputation uses the same prefix/decode split.

```python
import torch

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

mtq.quantize(model, recipe)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-5)
model.train()
optimizer.zero_grad(set_to_none=True)

# Two 128-token sequences: 64 and 96 prefill tokens, respectively.
# ids and labels have shape [2, 128] and are on the model's device.
with linear_attention_training_phase(model, [64, 96]):
with torch.autocast("cuda", dtype=torch.bfloat16):
loss = model(input_ids=ids, labels=labels, use_cache=False).loss
loss.backward()
optimizer.step()
```

For GDN, select `*gdn_state_quantizer` in the recipe instead of
`*kda_state_quantizer`. State quantization is dynamic, so these recipes use
`algorithm=None` without a calibration pass. Replay factors have their own FP8
toggle.

Policies persist through ModelOpt save/restore. Per-batch prefix lengths are
runtime metadata and must be supplied for each workload. The context restores
previous lengths on exit and supports nesting. Concurrent forwards on the same
model instance with different phase contexts are unsupported.

## Replay, decay, and Hadamard options

Token mode rounds each recurrent-state write. Replay mode retains an anchor and
ordered key/update factors, rounding the anchor at each `replay.window` refresh.
`factor_qdq` independently enables FP8 QDQ for those factors. `readout="working"`
reads the state before its write quantization; `"stored"` reads it afterward.

For KDA decay approximation, set `decode["decay_log_step"] = 1 / 256` before
conversion. It rounds suffix log retention before exponentiation with identity
STE gradients. Prefix decay remains exact.

The default INT8 codec applies a 32-point orthonormal Hadamard transform along
the value axis, quantizes token states or replay anchors in that basis, and
transforms outputs back. Value dimensions must be divisible by 32; `state.block_v` must be 32, 64,
or 128. Scales group one key channel and 32 values, independent of execution tile
width. INT8 codes use half-away-from-zero rounding with FP16 stored scales;
the optional tile codec instead uses nearest-even rounding and FP32 scales.

## Training boundaries

The Torch and Triton implementations support first-order QAT gradients through
initial states, chunk handoff, token writes, and replay refreshes. Triton supports
key dimensions up to 128 and value blocks 16/32/64/128, with checkpointed backward.
Use the Torch implementation for higher-order differentiation.

Prefill prefixes use exact chunk algebra with optional state QDQ. This example
has no prefill GEMM QDQ or approximate inverse. FLA KDA requires `use_cache=False`;
serving cache objects are rejected. ModelOpt saves execution policies, while
per-batch prefix lengths must be supplied again during training. Distributed
decode training and model-quality recovery require separate qualification.
Loading