Files
Model-Optimizer/tests/gpu/torch/export/test_export.py
T
Shengliang Xu d19925e446 simple refactor(export): split TensorRT-LLM-only code into modelopt/torch/export/trtllm (#2365)
### What does this PR do?

Type of change: refactor.

**The TensorRT-LLM checkpoint export format is deprecated.** Per
`docs/source/deployment/1_tensorrt_llm.rst`: *"The
`export_tensorrt_llm_checkpoint` API will be deprecated in future
releases. Users are encouraged to transition to the unified HF export
API, which provides enhanced functionality and flexibility for exporting
models to multiple inference frameworks including TensorRT-LLM, vLLM,
and SGLang."*

That deprecated code was not sitting off to one side — it was
**interleaved with the export path we actually want to grow.**
`modelopt/torch/export` mixed the deprecated TensorRT-LLM checkpoint
logic with the framework-agnostic HF/Megatron export code, in the same
modules:

- `layer_utils.py` was 1,986 lines, of which ~1,600 were TensorRT-LLM
`build_*_config` builders. The HF path imports this module for five
small predicates (`is_moe`, `is_quantlinear`, …) and dragged the whole
deprecated builder set in with them.
- `model_config.py` held the TensorRT-LLM `ModelConfig` dataclasses
*and* the `QUANTIZATION_*` / `KV_CACHE_*` constants that every backend
needs, so all of HF export imported the deprecated checkpoint schema to
get a format name string.
- `quant_utils.py` carried two helpers whose only caller is the
deprecated `postprocess.py`.

**This PR isolates the deprecated format so it stops polluting the
HuggingFace export path.** Everything reachable only from
`export_tensorrt_llm_checkpoint` now lives under
`modelopt/torch/export/trtllm/`, and the dependency is **one-way**:
`trtllm/` reaches into the parent through `quant_format`, `quant_utils`
and `layer_utils`, and **no implementation module in the parent imports
`trtllm/`.** The single exception is the deprecation re-export in
`modelopt/torch/export/__init__.py` described below, which is scheduled
for deletion in 0.49.0.

That one-way edge is the property worth protecting in review. It means
the deprecated format can be evolved, frozen, or eventually removed
without touching HF export, and HF export can no longer accidentally
grow a dependency on it.

### Deprecation handling

The format has carried a deprecation notice in the deployment docs since
`bc546943b4` (2025-10-08, first shipped in 0.39.0) — about 11 months.
But the deprecation policy in `README.md` also specifies *how* a
deprecation is communicated: a changelog entry, a source statement of
timing, and a runtime warning on use. **None of those existed**; only
one docs page ever said anything. So 0.48.0 is the first release that
gives users a signal they can act on, and this PR treats it as the
*start* of the migration period rather than the end:

- Both entry points now emit a `DeprecationWarning` naming 0.48.0 and
the 0.49.0 removal.
- `export_tensorrt_llm_checkpoint` and
`torch_to_tensorrt_llm_checkpoint` **remain importable from
`modelopt.torch.export`** for this release only, so existing callers
keep working *and* actually receive the warning. Removing the path in
the same release that first warns would mean callers hit `ImportError`
and never see it.
- The 0.49.0 removal date is stated in all four channels the policy
names: the runtime warning, the source (`.. deprecated:: 0.48.0` plus a
comment), the changelog, and the deployment doc.

The deeper module paths (`modelopt.torch.export.model_config_export`,
`modelopt.torch.export.model_config`) are **not** forwarded. Neither
appeared in a docs example, and `model_config.py` never declared
`__all__`, so by the `__all__` convention in `CONTRIBUTING.md` they were
never part of the public surface.

Eight modules had no non-TRT-LLM importer and moved whole:
`model_config_export`, `model_config_utils`, `postprocess`,
`distribute`, `tensorrt_llm_utils`, `tensorrt_llm_type`,
`hf_config_map`, `mcore_config_map`.

Three were genuinely mixed and were split by call-graph analysis rather
than by file:

| module | stayed shared (HF path) | moved to `trtllm/` (deprecated) |
|---|---|---|
| `model_config.py` | `QUANTIZATION_*`, `KV_CACHE_*`,
`FUSION_FREE_FORMATS` → new leaf module `quant_format.py` | the
`ModelConfig` dataclasses + `LINEAR_*`/`LAYERNORM_*` checkpoint-layout
constants |
| `layer_utils.py` | 9 module-shape predicates and MoE quantizer helpers
(`is_moe`, `is_quantlinear`, `get_experts_list`,
`sync_moe_gate_up_amax`, …) | the 39 `build_*_config` builders and
enc/dec helpers |
| `quant_utils.py` | everything else | `get_scaling_factor_from_weight`,
`resmooth_and_get_scale` (only caller is `trtllm/postprocess.py`) |

`adjust_attn_amax_values` was deliberately left in the shared
`quant_utils.py`: it has no production caller at all (only a test), so
"used only by TRT-LLM export" is not demonstrable for it.

Nothing was added or removed. `export_tensorrt_llm_checkpoint` behaves
exactly as before, just from a new import path and with a warning
attached.

### Usage

```python
# Deprecated TensorRT-LLM checkpoint export — new home, and warns on call
from modelopt.torch.export.trtllm import (
    export_tensorrt_llm_checkpoint,
    torch_to_tensorrt_llm_checkpoint,
)
from modelopt.torch.export.trtllm.model_config import ModelConfig

# The pre-0.48 path still works for one release, and warns — removed in 0.49.0
from modelopt.torch.export import export_tensorrt_llm_checkpoint

# Shared format constants — new home, still re-exported from the top level
from modelopt.torch.export.quant_format import QUANTIZATION_NVFP4, KV_CACHE_FP8
from modelopt.torch.export import QUANTIZATION_NVFP4  # still works

# The recommended path — unchanged
from modelopt.torch.export import export_hf_checkpoint, get_model_type
```

### Testing

- `pre-commit` on all changed files: passes (ruff, ruff-format,
**mypy**, bandit, markdownlint). mypy caught one implicit re-export of
`is_layernorm`, now imported from the shared module directly.
- `tests/unit/torch/export`: **189 passed**. With the new `trtllm/` test
dir: **193 passed**.
- Full `tests/unit/torch`: **2367 passed, 0 export failures**. The 45
failures are pre-existing environment issues — a deepspeed circular
import and a read-only HF cache — confirmed by reading their error text,
not assumed.
- `pytest tests/gpu/torch/export --collect-only`: 172 items, no
collection error.
- In-repo consumers updated and re-verified by an AST scan that imports
every `modelopt.torch.export*` module referenced anywhere in the tree
and checks each imported name still resolves: `hf_ptq.py`,
`export_trtllm_ckpt.py`, `deepseek_v3/ptq.py`, the AutoQuantize
notebook, `hf_ptq/README.md`, 2 docs pages, 4 tests.
- **Deprecation contract is covered by committed tests** (3 new, in the
`trtllm/` test dir): the pre-0.48 top-level import still resolves to the
same objects, `torch_to_tensorrt_llm_checkpoint` warns *at call time*
rather than on first `next()` (it returns a generator, so a naive
`warnings.warn` in the body would fire late or never), and one
`export_tensorrt_llm_checkpoint` call emits exactly one warning rather
than two. The first of these makes closing the migration window early a
test failure rather than a silent regression. `pyproject.toml` sets no
`filterwarnings = error`, so no suite fails on the new warning.
- **After merging `main`** (4 commits, incl. a 180-line rewrite of
`unified_export_megatron.py` that touches a file this PR also edits):
merged with no conflicts, then re-verified rather than trusted — import
scan clean across 24 export modules, `ruff check` clean repo-wide, 193
export tests passing, GPU collection still clean.

**Not run: the GPU suites** (`tests/gpu/torch/export`,
`tests/gpu_trtllm`) — no GPU in my environment.
`tests/gpu/torch/export/test_export.py` had its imports retargeted, so
it is the one most worth a GPU run before merge.

### Reviewer note: the deprecated path has no test coverage

Worth knowing before reviewing. **No test in the repo — including
`tests/examples/` — calls `export_tensorrt_llm_checkpoint`,
`torch_to_tensorrt_llm_checkpoint`, any `build_*_config`,
`convert_to_tensorrt_llm_config`, or `postprocess_model_config`.** So
~4,600 moved lines have no direct tests, and this refactor is validated
by import-graph reasoning, lint and mypy rather than by tests exercising
the moved code.

Given the format is deprecated and scheduled for removal in 0.49.0, **no
new coverage is planned for the conversion path itself** — writing fresh
tests for an API being removed next release isn't a good use of effort.
The gap is documented so reviewers can weigh the risk, not as a TODO.
(The deprecation *mechanism* is tested; see Testing.)

One caveat on how the gap was established: a runtime check showing all
12 `trtllm` modules in `sys.modules` after the export suites is *not*
evidence of coverage — importing any submodule runs
`trtllm/__init__.py`, which star-imports `model_config_export` and pulls
in the rest. Real line coverage could not be measured (`coverage`'s
tracer is incompatible with this venv's torch build: `ValueError: module
functions cannot set METH_CLASS or METH_STATIC`, on both the C tracer
and `sysmon`). The claim rests on a call-site audit generated from the
actual public symbols of those modules.

### Before your PR is "*Ready for review*"

- Is this change backward compatible?: ✅ for the public API —
`export_tensorrt_llm_checkpoint` and `torch_to_tensorrt_llm_checkpoint`
remain importable from `modelopt.torch.export` through the 0.49.0
migration period, now with a `DeprecationWarning`. The undocumented
submodule paths `modelopt.torch.export.model_config_export` and
`.model_config` did move; see **Usage**.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A — no new
code or dependencies; existing code relocated.
- Did you write any new necessary tests?: ✅ — 3 tests covering the
deprecation contract (old import path, call-time warning, exactly-one
warning). One existing test also moved to mirror the source split. No
new coverage for the deprecated conversion path itself; see the note
above.
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ — under 0.48.0 **Deprecations**, covering both the runtime warning and
the new import location.
- Did you get Claude approval on this PR?: ❌ — not yet run.

### Additional Information

Git detected the moves, so the diff stays reviewable: 8 files show as
pure renames (100%), the three split files as rename/copy at 94–99%
similarity, and only `layer_utils.py` as a 79% rewrite — expected, since
it shed 1,616 lines to `trtllm/`.

`examples/hf_ptq/hf_ptq.py` and
`examples/llm_sparsity/weight_sparsity/export_trtllm_ckpt.py` still call
the deprecated API, so those examples now print the warning. That is the
intended nudge, but happy to silence or migrate them if preferred. They
import from the new `.trtllm` path already, so they need no change at
0.49.0.

Two incidental changes, easy to revert if unwanted:
- `modelopt/torch/export/layer_utils.py` mode `100755 → 100644` (it was
needlessly executable).
- The new test is named `test_trtllm_quant_utils.py`, not
`test_quant_utils.py`: these directories have no `__init__.py`, so
pytest derives the module name from the bare filename and the shorter
name fails collection with `import file mismatch` against the existing
`test_quant_utils.py` one level up.


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

- **New Features**
- Added shared quantization and KV-cache format definitions for export
workflows.
- Added expanded TensorRT-LLM export support, including broader model
architecture and quantization handling.
- Added distributed export utilities for coordinating checkpoint data
across processes.

- **Deprecation**
- TensorRT-LLM checkpoint export now emits a warning and is scheduled
for removal in version 0.49.0.
- Use the documented export module and save optimized model state
explicitly when needed.

- **Documentation**
- Updated guides and examples with new import paths and deprecation
guidance.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Shengliang Xu <shengliangx@nvidia.com>
2026-09-10 10:44:04 -07:00

558 lines
20 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2024 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 json
from fnmatch import fnmatch
import pytest
import torch
from _test_utils.torch.export.utils import (
SmallLinearModelwithCustomWeight,
ToyModel,
only_input_quantizer_fp8_config,
only_output_quantizer_fp8_config,
only_weight_quantizer_fp8_config,
partial_fp8_config,
partial_fp8_kv_cache_config,
partial_int4_awq_config,
partial_int8_kv_cache_config,
partial_nvfp4_awq_config,
partial_nvfp4_config,
partial_w4a8_config,
)
from _test_utils.torch.transformers_models import get_tiny_qwen3_moe
import modelopt.torch.quantization as mtq
from modelopt.torch.export.quant_format import (
KV_CACHE_FP8,
KV_CACHE_INT8,
QUANTIZATION_FP8,
QUANTIZATION_INT4_AWQ,
QUANTIZATION_NONE,
QUANTIZATION_NVFP4,
QUANTIZATION_NVFP4_AWQ,
QUANTIZATION_W4A8_AWQ,
)
from modelopt.torch.export.quant_utils import (
adjust_attn_amax_values,
all_items_same,
get_kv_cache_dtype,
get_quant_config,
get_quantization_format,
get_scaling_factor,
get_weight_block_size,
postprocess_state_dict,
process_layer_quant_config,
to_quantized_weight,
)
from modelopt.torch.export.unified_export_hf import export_hf_checkpoint
from modelopt.torch.quantization.config import (
FP8_DEFAULT_CFG,
INT4_AWQ_CFG,
INT8_SMOOTHQUANT_CFG,
INT8_WEIGHT_ONLY_CFG,
NVFP4_AWQ_LITE_CFG,
NVFP4_DEFAULT_CFG,
NVFP4_EXPERTS_ONLY_CFG,
W4A8_AWQ_BETA_CFG,
)
from modelopt.torch.quantization.nn import SequentialQuantizer, TensorQuantizer
from modelopt.torch.quantization.qtensor import INT4QTensor, QTensorWrapper
@pytest.mark.parametrize(
("config", "expected"),
[
(partial_fp8_config, QUANTIZATION_FP8),
(partial_w4a8_config, QUANTIZATION_W4A8_AWQ),
(partial_nvfp4_config, QUANTIZATION_NVFP4),
(partial_nvfp4_awq_config, QUANTIZATION_NVFP4_AWQ),
(partial_int4_awq_config, QUANTIZATION_INT4_AWQ),
],
)
def test_get_quantization_format(config, expected):
model = ToyModel().to("cuda")
mtq.quantize(model, config, lambda x: x(torch.randn(1, 4, 10, device="cuda")))
assert get_quantization_format(model) == expected
@pytest.mark.parametrize(
("layer_config_dict", "expected_processed_dict"),
[
(
{
"layer1.quantization": "nvfp4", # All qformats
"layer1.awq_block_size": 16,
"layer3.quantization": "int4_awq",
"layer3.awq_block_size": 8,
"layer4.quantization": "w4a8_awq",
"layer4.awq_block_size": 64,
"layer5.quantization": "int8_sq",
"layer6.quantization": "fp8",
"layer7.quantization": "xyz",
"layer8.quantization": None,
},
{
"quant_algo": "MIXED_PRECISION",
"kv_cache_quant_algo": None,
"quantized_layers": {
"layer1": {"quant_algo": "NVFP4", "group_size": 16},
"layer3": {
"quant_algo": "W4A16_AWQ",
"group_size": 8,
"has_zero_point": False,
"pre_quant_scale": True,
},
"layer4": {
"quant_algo": "W4A8_AWQ",
"group_size": 64,
"has_zero_point": False,
"pre_quant_scale": True,
},
"layer5": {"quant_algo": "W8A8_SQ_PER_CHANNEL"},
"layer6": {"quant_algo": "FP8"},
"layer7": {"quant_algo": "xyz"},
},
},
),
(
{
"layer1.quantization": "nvfp4", # Auto quant with one qformat case
"layer1.awq_block_size": 16,
"layer2.quantization": "nvfp4",
"layer2.awq_block_size": 16,
"layer8.quantization": None,
},
{
"quant_algo": "NVFP4",
"kv_cache_quant_algo": None,
"group_size": 16,
"exclude_modules": ["layer8"],
},
),
],
)
def test_process_layer_quant_config(layer_config_dict, expected_processed_dict):
per_layer_config = process_layer_quant_config(layer_config_dict)
assert per_layer_config == expected_processed_dict
@pytest.mark.parametrize(
("item_list", "expected"),
[
(["a", "a", "a"], True),
(["b", "a", "a"], False),
(["a"], True),
([True, True, True], True),
([True, False, True], False),
([False, False, False], True),
([False], True),
([True], True),
],
)
def test_all_items_same(item_list, expected):
generated = all_items_same(item_list)
assert generated == expected
@pytest.mark.parametrize(
("state_dict", "quantization", "maxbound", "expected_state_dict"),
[
( # Test replacements and KV cache scaling
{
"layer1.k_bmm_quantizer._amax": torch.tensor([0.128]),
"layer1.v_bmm_quantizer._amax": torch.tensor([256.0]),
"layer1.input_quantizer._pre_quant_scale": torch.tensor([0.128]),
},
KV_CACHE_FP8,
128.0,
{
"layer1.k_proj.k_scale": torch.tensor([0.001]),
"layer1.v_proj.v_scale": torch.tensor([2.0]),
"layer1.pre_quant_scale": torch.tensor([0.128]),
},
),
( # Test skipping output_quantizer _amax keys other than k_scale and v_scale
{
"layer1.q_bmm_quantizer._amax": torch.tensor([0.128]),
"layer1.k_bmm_quantizer._amax": torch.tensor([0.128]),
"layer1.v_bmm_quantizer._amax": torch.tensor([256]),
},
KV_CACHE_FP8,
128.0,
{
"layer1.k_proj.k_scale": torch.tensor([0.001]),
"layer1.v_proj.v_scale": torch.tensor([2.0]),
},
),
( # Test squeezing tensor with leading dimension 1
{
"layer1.k_proj.weight": torch.ones(1, 1),
},
KV_CACHE_FP8,
128.0,
{
"layer1.k_proj.weight": torch.ones(1),
},
),
( # Test case with no KV cache scaling + AWQ quant
{
"layer1.input_quantizer._pre_quant_scale": torch.tensor([0.128]),
},
QUANTIZATION_NONE,
128.0,
{
"layer1.pre_quant_scale": torch.tensor([0.128]),
},
),
],
)
def test_postprocess_state_dict(state_dict, quantization, maxbound, expected_state_dict):
processed_state_dict = postprocess_state_dict(state_dict, maxbound, quantization)
assert processed_state_dict == expected_state_dict
def test_postprocess_state_dict_qlora_strips_base_layer():
"""Every QLoRA `base_layer.*` tensor needed for deployment must survive the rename.
Dropping the NVFP4 global scale or a bias yields an undeployable checkpoint.
"""
state_dict = {
"layer1.base_layer.weight": torch.ones(4, 2, dtype=torch.uint8),
"layer1.base_layer.weight_scale": torch.ones(4, 1),
"layer1.base_layer.weight_scale_2": torch.tensor([0.5]),
"layer1.base_layer.input_scale": torch.tensor([0.25]),
"layer1.base_layer.bias": torch.arange(4.0),
"layer1.base_layer.input_quantizer._pre_quant_scale": torch.ones(2),
# Quantizer internals must still be dropped.
"layer1.base_layer.weight_quantizer._amax": torch.tensor([1.0]),
"layer1.base_layer.input_quantizer._amax": torch.tensor([1.0]),
"layer1.base_layer.weight_quantizer._scale": torch.ones(4, 1),
"layer1.base_layer.weight_quantizer._double_scale": torch.tensor([0.5]),
}
processed_state_dict = postprocess_state_dict(
state_dict, 448.0, QUANTIZATION_NONE, is_modelopt_qlora=True
)
assert set(processed_state_dict) == {
"layer1.weight",
"layer1.weight_scale",
"layer1.weight_scale_2",
"layer1.input_scale",
"layer1.bias",
"layer1.pre_quant_scale",
}
assert torch.equal(processed_state_dict["layer1.weight_scale_2"], torch.tensor([0.5]))
assert torch.equal(processed_state_dict["layer1.bias"], torch.arange(4.0))
@pytest.mark.parametrize(
("config", "expected"),
[
(partial_fp8_kv_cache_config, KV_CACHE_FP8),
(partial_int8_kv_cache_config, KV_CACHE_INT8),
],
)
def test_get_kv_cache_dtype(config, expected):
model = ToyModel().to("cuda")
mtq.quantize(model, config, lambda x: x(torch.randn(1, 4, 10, device="cuda")))
# Create list of modules in model
modules = []
for name, module in model.named_modules():
modules.append(module)
kv_cache_dtype = get_kv_cache_dtype(modules)
assert kv_cache_dtype == expected
# Tensor Quantizer extraction for export tests
@pytest.mark.parametrize(
("q_weight", "k_weight", "v_weight", "o_weight", "expected_qkv_amax", "expected_o_amax"),
[
(
torch.tensor([[0.1, 0.3], [0.22, 0.45]]),
torch.tensor([[0.44, 0.32], [0.11, 0.95]]),
torch.tensor([[0.9, 0.03], [0.92, 0.8]]),
torch.tensor([[0.01, 0.97], [0.29, 0.77]]),
0.95,
0.97,
),
(
torch.tensor([[0.1, 0.3], [0.22, 0.45]]),
torch.tensor([[0.44, 0.32], [0.11, -0.95]]),
torch.tensor([[0.9, 0.03], [0.92, 0.8]]),
torch.tensor([[0.01, -0.97], [0.29, 0.77]]),
0.95,
0.97,
),
],
)
@pytest.mark.parametrize(
"config",
[
FP8_DEFAULT_CFG,
NVFP4_DEFAULT_CFG,
],
)
def test_adjust_attn_amax_values(
q_weight, k_weight, v_weight, o_weight, expected_qkv_amax, expected_o_amax, config
):
# Initialize model and quantize to insert quantizers
model = SmallLinearModelwithCustomWeight([q_weight, k_weight, v_weight, o_weight]).to("cuda")
mtq.quantize(model, config, lambda x: x(torch.randn(1, 4, q_weight.shape[1], device="cuda")))
adjust_attn_amax_values(model)
# Weight quantizer amax must remain unchanged for non qkv layers
assert (
model.q_proj.weight_quantizer.amax
== model.k_proj.weight_quantizer.amax
== model.v_proj.weight_quantizer.amax
== expected_qkv_amax
)
# Weight quantizer amax must be updated for q,k,v layers
assert model.o_proj.weight_quantizer.amax == expected_o_amax
@pytest.mark.parametrize(
("config", "expected_block_size"),
[
(FP8_DEFAULT_CFG, 0),
(INT8_WEIGHT_ONLY_CFG, 0),
(INT8_SMOOTHQUANT_CFG, 0),
(NVFP4_DEFAULT_CFG, 16),
(NVFP4_AWQ_LITE_CFG, 16),
(W4A8_AWQ_BETA_CFG, 128),
(INT4_AWQ_CFG, 128),
(partial_nvfp4_config, 16),
],
)
def test_get_weight_block_size(config, expected_block_size):
model = ToyModel().to("cuda")
mtq.quantize(model, config, lambda x: x(torch.randn(1, 4, 10, device="cuda")))
for _, module in model.named_modules():
block_size = get_weight_block_size(module)
if hasattr(module, "weight_quantizer"):
if (
isinstance(module.weight_quantizer, SequentialQuantizer)
or module.weight_quantizer.is_enabled
):
assert block_size == expected_block_size
else:
assert block_size == 0
else:
assert block_size == 0
@pytest.mark.parametrize("quantization", [QUANTIZATION_INT4_AWQ, QUANTIZATION_W4A8_AWQ])
def test_to_quantized_weight_int4_block_size(quantization):
block_size = 128
in_dim = 2 * block_size
scales = torch.tensor([[1.0, 2.0]] * 4, device="cuda")
quantized_values = torch.arange(1, 5, device="cuda")[:, None]
weight = scales.repeat_interleave(block_size, dim=-1) * quantized_values
packed = to_quantized_weight(weight, scales, quantization, block_size=block_size)
assert packed.shape == (2, in_dim)
assert torch.equal(packed[0], torch.full((in_dim,), 0x21, dtype=torch.uint8, device="cuda"))
assert torch.equal(packed[1], torch.full((in_dim,), 0x43, dtype=torch.uint8, device="cuda"))
partial_weight = torch.cat((weight, quantized_values.repeat(1, 2)), dim=-1)
with pytest.raises(NotImplementedError, match="partial blocks are not supported"):
to_quantized_weight(partial_weight, scales, quantization, block_size=block_size)
compressed_weight, _ = INT4QTensor.quantize(partial_weight, block_size)
with pytest.raises(NotImplementedError, match="partial blocks are not supported"):
to_quantized_weight(
QTensorWrapper(compressed_weight), scales, quantization, block_size=block_size
)
@pytest.mark.parametrize("quantization", [QUANTIZATION_INT4_AWQ, QUANTIZATION_W4A8_AWQ])
@pytest.mark.parametrize("block_size", [None, 0, -1, 2.0])
def test_to_quantized_weight_invalid_int4_block_size(quantization, block_size):
weight = torch.ones((4, 4), device="cuda")
scales = torch.ones((4, 2), device="cuda")
with pytest.raises(ValueError, match="Block size must be a positive integer"):
to_quantized_weight(weight, scales, quantization, block_size=block_size)
@pytest.mark.parametrize(
("config", "maxbound", "expected_amax"),
[
(only_weight_quantizer_fp8_config, 448, [0.45, 0.95, 0.92, 0.97]),
(only_input_quantizer_fp8_config, 448, [1.0, 0.67, 0.68, 0.9]),
(only_output_quantizer_fp8_config, 448, [0.67, 0.68, 0.9, 0.88]),
],
)
@pytest.mark.parametrize(
("q_weight", "k_weight", "v_weight", "o_weight"),
[
(
torch.tensor([[0.1, 0.3], [0.22, 0.45]]),
torch.tensor([[0.44, 0.32], [0.11, 0.95]]),
torch.tensor([[0.9, 0.03], [0.92, 0.8]]),
torch.tensor([[0.01, 0.97], [0.29, 0.77]]),
),
],
)
def test_get_scaling_factor(
q_weight, k_weight, v_weight, o_weight, config, expected_amax, maxbound
):
# Initialize model and quantize to insert quantizers
model = SmallLinearModelwithCustomWeight([q_weight, k_weight, v_weight, o_weight]).to("cuda")
mtq.quantize(model, config, lambda x: x(torch.ones(1, 2, q_weight.shape[1], device="cuda")))
for name, module in model.named_modules():
if isinstance(module, TensorQuantizer) and module.is_enabled:
scale = get_scaling_factor(module)
print(f"DEBUG LOG: Scale: {scale}, Expected: {expected_amax[0] / maxbound}")
assert torch.allclose(
scale,
torch.tensor((expected_amax[0] / maxbound), dtype=scale.dtype),
rtol=1e-3,
atol=1e-3,
)
expected_amax.pop(0)
@pytest.mark.parametrize(
("config", "expected"),
[
(
partial_fp8_config,
{
"exclude_modules": ["linears.0", "linears.2"],
"kv_cache_quant_algo": None,
"quant_algo": "FP8",
},
),
(
partial_w4a8_config,
{
"exclude_modules": ["linears.0", "linears.1"],
"group_size": 128,
"has_zero_point": False,
"kv_cache_quant_algo": None,
"pre_quant_scale": True,
"quant_algo": "W4A8_AWQ",
},
),
(
partial_nvfp4_config,
{
"exclude_modules": ["linears.0"],
"group_size": 16,
"kv_cache_quant_algo": None,
"quant_algo": "NVFP4",
},
),
(
partial_nvfp4_awq_config,
{
"exclude_modules": ["linears.0", "linears.1"],
"group_size": 16,
"has_zero_point": False,
"pre_quant_scale": True,
"kv_cache_quant_algo": None,
"quant_algo": "NVFP4_AWQ",
},
),
(
partial_int4_awq_config,
{
"exclude_modules": ["linears.0", "linears.1"],
"group_size": 128,
"has_zero_point": False,
"kv_cache_quant_algo": None,
"pre_quant_scale": True,
"quant_algo": "W4A16_AWQ",
},
),
(
partial_fp8_kv_cache_config,
{
"exclude_modules": ["linears.0", "linears.2"],
"quant_algo": "FP8",
"kv_cache_quant_algo": "FP8",
},
),
(
partial_int8_kv_cache_config,
{
"exclude_modules": ["linears.0", "linears.2"],
"quant_algo": "FP8",
"kv_cache_quant_algo": "INT8",
},
),
],
)
def test_get_quant_config(config, expected):
model = ToyModel().to("cuda")
mtq.quantize(model, config, lambda x: x(torch.randn(1, 4, 10, device="cuda")))
quant_config = get_quant_config(model)
assert quant_config["quantization"] == expected
def test_qwen3_moe_nvfp4_experts_only_export_exclude_modules(tmp_path):
"""Test that NVFP4_EXPERTS_ONLY_CFG correctly excludes non-expert modules in HF export.
For a Qwen3 MoE model, only routed expert layers (mlp.experts.*) should be quantized.
Attention layers and lm_head should appear in the exported hf_quant_config.json
exclude_modules.
Reference: https://huggingface.co/nvidia/Qwen3.5-397B-A17B-NVFP4/blob/main/hf_quant_config.json
"""
model = get_tiny_qwen3_moe().to("cuda")
# from_config doesn't set architectures; export code requires it
model.config.architectures = ["Qwen3MoeForCausalLM"]
# Quantize with NVFP4_EXPERTS_ONLY_CFG (targets only *mlp.experts* patterns)
dummy_inputs = {k: v.to("cuda") for k, v in model.dummy_inputs.items()}
mtq.quantize(model, NVFP4_EXPERTS_ONLY_CFG, lambda m: m(**dummy_inputs))
# Export
export_dir = tmp_path / "qwen3_moe_nvfp4_experts_only"
export_hf_checkpoint(model, export_dir=export_dir)
# Load the generated hf_quant_config.json
hf_quant_config_path = export_dir / "hf_quant_config.json"
assert hf_quant_config_path.exists(), "hf_quant_config.json should be generated"
with open(hf_quant_config_path) as f:
hf_quant_config = json.load(f)
quant_section = hf_quant_config["quantization"]
assert quant_section["quant_algo"] == "NVFP4"
exclude_modules = quant_section["exclude_modules"]
def is_excluded(module_name: str) -> bool:
return any(fnmatch(module_name, pattern) for pattern in exclude_modules)
# Attention layers must be excluded
assert is_excluded("model.layers.0.self_attn.q_proj"), (
f"self_attn should be excluded, got patterns: {exclude_modules}"
)
# lm_head must be excluded
assert is_excluded("lm_head"), f"lm_head should be excluded, got patterns: {exclude_modules}"
# Routed experts should NOT be excluded
assert not is_excluded("model.layers.0.mlp.experts.0.down_proj"), (
f"Routed experts should not be excluded, got patterns: {exclude_modules}"
)