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>
This commit is contained in:
Shengliang Xu
2026-09-10 10:44:04 -07:00
committed by GitHub
parent 28dc117594
commit d19925e446
33 changed files with 2078 additions and 1796 deletions
+2
View File
@@ -14,6 +14,8 @@ Changelog
**Deprecations**
- The TensorRT-LLM checkpoint export format is deprecated and will be removed in 0.49.0: ``export_tensorrt_llm_checkpoint`` and ``torch_to_tensorrt_llm_checkpoint`` now emit a ``DeprecationWarning`` on use. Use ``export_hf_checkpoint``, which exports a unified Hugging Face checkpoint deployable on TensorRT-LLM, vLLM and SGLang. Its implementation moved to ``modelopt.torch.export.trtllm``, so import those two functions from there and the ``ModelConfig`` dataclasses from ``modelopt.torch.export.trtllm.model_config``; both functions remain importable from ``modelopt.torch.export`` for this release only.
**Bug Fixes**
- Fix ``--use_fsdp2`` HuggingFace checkpoint export gathering the whole model onto rank 0, which made export the dominant phase of a PTQ run and could exhaust host memory on large models. The model is now split into per-decoder-layer units dealt round-robin across ranks; each rank gathers every unit but keeps, packs, and writes only the ones it owns, so a rank buffers roughly ``model / world_size`` instead of the whole checkpoint, and rank 0 writes the combined index. Export configurations that cannot be split this way now raise instead of producing a mismatched checkpoint: FSDP2 combined with another DTensor parallelism (for example FSDP2 + tensor parallel on a 2-D mesh; HSDP is supported), models whose decoder layers cannot be discovered, a decoder layer object reused across layers, and a module that holds the decoder layers while owning parameters of its own.
+4 -4
View File
@@ -2,7 +2,7 @@
TensorRT-LLM
==========================
**Deprecation Notice**: The export_tensorrt_llm_checkpoint API will be deprecated in future releases. Users are encouraged to transition to the :doc:`unified HF export API <3_unified_hf>`, which provides enhanced functionality and flexibility for exporting models to multiple inference frameworks including TensorRT-LLM, vLLM, and SGLang.
**Deprecation Notice**: The export_tensorrt_llm_checkpoint API is deprecated as of 0.48.0 and will be removed in 0.49.0. Users are encouraged to transition to the :doc:`unified HF export API <3_unified_hf>`, which provides enhanced functionality and flexibility for exporting models to multiple inference frameworks including TensorRT-LLM, vLLM, and SGLang.
.. note::
@@ -27,11 +27,11 @@ After the model is quantized, the quantized model can be exported to the TensorR
#. A single JSON file recording the model structure and metadata (config.json)
#. A group of safetensors files, each recording the local calibrated model on a single GPU rank (model weights, scaling factors per GPU).
The export API (:meth:`export_tensorrt_llm_checkpoint <modelopt.torch.export.model_config_export.export_tensorrt_llm_checkpoint>`) can be used as follows:
The export API (:meth:`export_tensorrt_llm_checkpoint <modelopt.torch.export.trtllm.model_config_export.export_tensorrt_llm_checkpoint>`) can be used as follows:
.. code-block:: python
from modelopt.torch.export import export_tensorrt_llm_checkpoint
from modelopt.torch.export.trtllm import export_tensorrt_llm_checkpoint
with torch.inference_mode():
export_tensorrt_llm_checkpoint(
@@ -43,7 +43,7 @@ The export API (:meth:`export_tensorrt_llm_checkpoint <modelopt.torch.export.mod
inference_pipeline_parallel, # The number of GPUs used in the inference time for pipeline parallelism.
)
If the :meth:`export_tensorrt_llm_checkpoint <modelopt.torch.export.model_config_export.export_tensorrt_llm_checkpoint>` call is successful, the TensorRT-LLM checkpoint will be saved. Otherwise, e.g. the ``decoder_type`` is not supported, a torch state_dict checkpoint will be saved instead.
If the :meth:`export_tensorrt_llm_checkpoint <modelopt.torch.export.trtllm.model_config_export.export_tensorrt_llm_checkpoint>` call is successful, the TensorRT-LLM checkpoint will be saved. Otherwise, e.g. the ``decoder_type`` is not supported, the call warns and re-raises the exception, and no checkpoint is written. To inspect the model in that case, save the ModelOpt-optimized ``state_dict`` yourself with ``torch.save``.
.. list-table:: Model support matrix for the TensorRT-LLM checkpoint export
:header-rows: 1
@@ -16,7 +16,7 @@ As ModelOpt cannot detect these linear ops out-of-the-box, a HugggingFace plugin
#. Rewrite the linear ops (w1, v1 and v2) as a standard ``nn.Linear`` op, and re-implement the ``forward`` method.
#. Register the new dynamic ``_QuantDbrxExperts`` to replace the ``DbrxExperts`` from the modeling_dbrx.py in the ``transformers`` library
#. Try quantize the DBRX model after the plugin is implemented, feel free to follow the `hf_ptq example <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/hf_ptq>`_.
#. TensorRT-LLM is open-sourced. If this customized model is not supported by TensorRT-LLM yet, please modify :meth:`export_tensorrt_llm_checkpoint <modelopt.torch.export.export_tensorrt_llm_checkpoint>` or :meth:`export_hf_checkpoint <modelopt.torch.export.export_hf_checkpoint>` to export the quantized model for deployment with a customized TensorRT-LLM modeling implementation. Feel free to :doc:`contact us <../support/1_contact>` if further support is needed.
#. TensorRT-LLM is open-sourced. If this customized model is not supported by TensorRT-LLM yet, please modify :meth:`export_tensorrt_llm_checkpoint <modelopt.torch.export.trtllm.export_tensorrt_llm_checkpoint>` or :meth:`export_hf_checkpoint <modelopt.torch.export.export_hf_checkpoint>` to export the quantized model for deployment with a customized TensorRT-LLM modeling implementation. Feel free to :doc:`contact us <../support/1_contact>` if further support is needed.
The following code snippet is excerpted from ``modelopt/torch/quantization/plugins/huggingface.py``
+1 -1
View File
@@ -53,7 +53,7 @@ from tqdm import tqdm
from transformers import AutoTokenizer
import modelopt.torch.quantization as mtq
from modelopt.torch.export.model_config import KV_CACHE_FP8
from modelopt.torch.export.quant_format import KV_CACHE_FP8
from modelopt.torch.export.quant_utils import get_quant_config
from modelopt.torch.kernels.quantization.gemm import weight_dequant
from modelopt.torch.quantization.nn import TensorQuantizer
+1 -1
View File
@@ -653,7 +653,7 @@ version behind each entry.
The user can specify the inference time TP and PP size and the export API will organize the weights to fit the target GPUs.
```python
from modelopt.torch.export import export_tensorrt_llm_checkpoint
from modelopt.torch.export.trtllm import export_tensorrt_llm_checkpoint
with torch.inference_mode():
export_tensorrt_llm_checkpoint(
+1 -1
View File
@@ -76,13 +76,13 @@ from modelopt.torch.export import (
export_hf_checkpoint,
export_hf_vllm_fq_checkpoint,
export_speculative_decoding,
export_tensorrt_llm_checkpoint,
get_model_type,
has_spec_opt,
save_expert_token_count_table,
)
from modelopt.torch.export.layerwise_export import LayerwiseExporter
from modelopt.torch.export.model_utils import get_language_model_from_vl, is_multimodal_model
from modelopt.torch.export.trtllm import export_tensorrt_llm_checkpoint
from modelopt.torch.quantization.config import need_calibration
from modelopt.torch.quantization.plugins.accelerate import init_quantized_weights
from modelopt.torch.quantization.utils import is_quantized
@@ -366,7 +366,8 @@
"metadata": {},
"outputs": [],
"source": [
"from modelopt.torch.export import export_hf_checkpoint, export_tensorrt_llm_checkpoint\n",
"from modelopt.torch.export import export_hf_checkpoint\n",
"from modelopt.torch.export.trtllm import export_tensorrt_llm_checkpoint\n",
"\n",
"if EXPORT_FMT == \"tensorrt_llm\":\n",
" export_tensorrt_llm_checkpoint(\n",
@@ -23,7 +23,8 @@ from transformers import AutoModelForCausalLM, AutoTokenizer, PreTrainedModel, P
import modelopt.torch.opt as mto
import modelopt.torch.sparsity as mts
from modelopt.torch.export import export_tensorrt_llm_checkpoint, get_model_type
from modelopt.torch.export import get_model_type
from modelopt.torch.export.trtllm import export_tensorrt_llm_checkpoint
DEFAULT_PAD_TOKEN = "[PAD]"
+11 -2
View File
@@ -16,13 +16,22 @@
"""Export package for Hugging Face and Megatron-based models."""
from .convert_hf_config import *
from .model_config import *
from .model_config_export import *
from .model_utils import *
from .moe_utils import *
from .plugins import *
from .quant_format import *
from .registry import *
from .shard_cast_utils import *
from .transformer_engine import *
# Deprecated: kept only to satisfy the migration period in the deprecation policy (README.md),
# which requires a deprecated feature to keep working while warning for one release. The
# TensorRT-LLM checkpoint export moved to ``modelopt.torch.export.trtllm`` in 0.48.0; these two
# names are its previously documented import path. Both warn on call. Remove this re-export in
# 0.49.0 -- nothing inside this package may depend on it.
from .trtllm.model_config_export import (
export_tensorrt_llm_checkpoint,
torch_to_tensorrt_llm_checkpoint,
)
from .unified_export_hf import *
from .unified_export_megatron import *
+1 -1
View File
@@ -21,8 +21,8 @@ import warnings
import torch.nn as nn
from .layer_utils import get_expert_linear_names, is_quantlinear, set_expert_quantizer_amax
from .model_config import QUANTIZATION_NONE
from .moe_utils import _export_fused_experts
from .quant_format import QUANTIZATION_NONE
from .quant_utils import get_quantization_format
from .registry import ExportContext, ExportModuleRegistry, PrepareMoEInputsRegistry
+3 -1595
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -35,9 +35,9 @@ from modelopt.torch.quantization.utils.layerwise_calib import LayerActivationCol
from modelopt.torch.utils import distributed as dist
from .layer_utils import is_moe, sync_moe_gate_up_amax
from .model_config import FUSION_FREE_FORMATS, QUANTIZATION_NVFP4
from .model_utils import TiedWeightMap, get_language_model_from_vl
from .quant_aware_conversion import build_reverse_name_mapper, revert_quant_config_names
from .quant_format import FUSION_FREE_FORMATS, QUANTIZATION_NVFP4
from .quant_utils import _postprocess_single_tensor, get_quant_config, get_quantization_format
from .registry import ExportContext, PrepareMoEInputsRegistry
from .unified_export_hf import (
@@ -22,7 +22,7 @@ from typing import Any
import torch
from modelopt.torch.export.model_config import QUANTIZATION_NONE
from modelopt.torch.export.quant_format import QUANTIZATION_NONE
from modelopt.torch.export.unified_export_megatron import GPTModelExporter
from modelopt.torch.quantization.utils import get_quantizer_state_dict
from modelopt.torch.utils.distributed import DistributedProcessGroup, is_master
+48
View File
@@ -0,0 +1,48 @@
# 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.
"""The quantization and KV-cache format names shared by every export backend.
Backend-specific names live with their backend: the TensorRT-LLM checkpoint layout
constants, for example, are in :mod:`modelopt.torch.export.trtllm.model_config`.
"""
QUANTIZATION_NONE = None
QUANTIZATION_FP8 = "fp8"
QUANTIZATION_INT8_SQ = "int8_sq"
QUANTIZATION_INT8_WO = "int8_wo"
QUANTIZATION_INT4_AWQ = "int4_awq"
QUANTIZATION_W4A8_AWQ = "w4a8_awq"
QUANTIZATION_NVFP4 = "nvfp4"
QUANTIZATION_NVFP4_SVDQUANT = "nvfp4_svdquant"
QUANTIZATION_W4A8_NVFP4_FP8 = "w4a8_nvfp4_fp8"
QUANTIZATION_MXFP4 = "mxfp4"
QUANTIZATION_MXFP8 = "mxfp8"
QUANTIZATION_W4A8_MXFP4_FP8 = "w4a8_mxfp4_fp8"
QUANTIZATION_W4A16_NVFP4 = "w4a16_nvfp4"
QUANTIZATION_NVFP4_AWQ = "nvfp4_awq"
QUANTIZATION_FP8_PB_REAL = "fp8_pb_real"
QUANTIZATION_FP8_PB_WO = "fp8_pb_wo"
QUANTIZATION_FP8_PC_PT = "fp8_pc_pt"
# Formats whose scales are purely per-module, so export never merges them across the q/k/v
# and gate/up groups that share an input. Every other format unifies input_amax (and, for
# NVFP4, weight_scale_2) across such a group, which only a whole-model forward can discover.
FUSION_FREE_FORMATS = frozenset({QUANTIZATION_FP8, QUANTIZATION_NONE, QUANTIZATION_FP8_PB_REAL})
KV_CACHE_FP8 = "FP8"
KV_CACHE_INT8 = "INT8"
KV_CACHE_NVFP4 = "NVFP4"
KV_CACHE_NVFP4_AFFINE = "NVFP4_AFFINE"
+2 -87
View File
@@ -48,7 +48,8 @@ from modelopt.torch.quantization.utils import (
from modelopt.torch.utils import clear_cuda_cache
from ..quantization.nn import NVFP4StaticQuantizer, SequentialQuantizer, TensorQuantizer
from .model_config import (
from .model_utils import TiedWeightMap
from .quant_format import (
KV_CACHE_FP8,
KV_CACHE_INT8,
KV_CACHE_NVFP4,
@@ -71,37 +72,10 @@ from .model_config import (
QUANTIZATION_W4A8_NVFP4_FP8,
QUANTIZATION_W4A16_NVFP4,
)
from .model_utils import TiedWeightMap
logger = logging.getLogger(__name__)
def get_scaling_factor_from_weight(weight, group_size) -> torch.tensor:
"""Calculate the weight scaling factor for a given group size."""
[n, k] = weight.shape
if group_size != 0:
# int4_awq
if k % group_size != 0:
raise NotImplementedError(
"Weight shape is not divisible for block size for block quantization."
)
weight = weight.reshape(n, k // group_size, group_size)
maxbound = 7.0
else:
# int8_sq
maxbound = 127.0
amax = weight.abs().max(dim=-1)[0].float()
weights_scaling_factor = amax / maxbound
# Let's filter the zeros in the scaling factor if the weights are zero
# to avoid the divided-by-zero error..
weights_scaling_factor[weights_scaling_factor == 0] = 1.0
return weights_scaling_factor
def maybe_transpose_expert_weight_dimensions(
weight: torch.Tensor,
weight_scale: torch.Tensor | None = None,
@@ -134,65 +108,6 @@ def maybe_transpose_expert_weight_dimensions(
return transposed_weight, transposed_weight_scale
def resmooth_and_get_scale(
merged_weights: torch.Tensor,
pre_quant_scales: list[torch.Tensor],
ranks: int,
group_size: int,
new_pre_quant_scale: torch.Tensor | None = None,
quantization: str | None = QUANTIZATION_NONE,
):
"""Resmooths weights from a single or multiple ranks and get scaling factors and amax.
Args:
merged_weights: Merged weights from ranks.
pre_quant_scales: List of pre-quantization scales for each rank.
ranks: Number of ranks.
group_size: Group size of the quantization block.
new_pre_quant_scale (optional): If not provided, weights will be resmoothed using
the average of pre_quant_scales.
Returns:
weights: Resmoothed weights.
weight_scaling_factors: Resmoothed scaling factors.
avg_pre_quant_scale: Calculated average of the quantization scale.
"""
if new_pre_quant_scale is None:
new_pre_quant_scale = torch.stack(pre_quant_scales).mean(dim=0)
assert len(pre_quant_scales) > 0 and new_pre_quant_scale.numel() == merged_weights.shape[1], (
"Shape of pre_quant_scales and weights do not match."
)
weights = torch.chunk(merged_weights, ranks, dim=0)
scales = []
new_weights = []
for i, p_scaling_factor in enumerate(pre_quant_scales):
# De smooth & Re smooth
weight = (
weights[i]
* p_scaling_factor.type(weights[i].dtype)
/ new_pre_quant_scale.type(weights[i].dtype)
)
new_weights.append(weight)
# If NVFP4_AWQ then we view the scales as uint8 to allow for cat later
if quantization in [QUANTIZATION_NVFP4_AWQ, QUANTIZATION_NVFP4_SVDQUANT]:
scale, _ = NVFP4QTensor.get_weights_scaling_factor(weight, group_size).view(torch.uint8)
else:
scale = get_scaling_factor_from_weight(weight, group_size)
scales.append(scale)
resmoothed_scales = torch.cat(scales, dim=0)
return (
torch.cat(new_weights, dim=0),
resmoothed_scales.view(torch.float8_e4m3fn)
if quantization in [QUANTIZATION_NVFP4_AWQ, QUANTIZATION_NVFP4_SVDQUANT]
else resmoothed_scales, # if NVFP4_AWQ we view the scales back as float8_e4m3fn after cat
new_pre_quant_scale,
)
def adjust_attn_amax_values(module):
"""Adjusts the amax values for the attention layers."""
projection_prefixes = ["q", "k", "v"]
+26
View File
@@ -0,0 +1,26 @@
# 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.
"""Checkpoint export logic for the TensorRT-LLM specific format.
**Deprecation Notice**: The ``export_tensorrt_llm_checkpoint`` API is deprecated as of 0.48.0 and
will be removed in 0.49.0. Users are encouraged to transition to the unified HF export API
(:meth:`export_hf_checkpoint <modelopt.torch.export.unified_export_hf.export_hf_checkpoint>`),
which provides enhanced functionality and flexibility for exporting models to multiple inference
frameworks including TensorRT-LLM, vLLM, and SGLang.
"""
from .model_config import *
from .model_config_export import *
File diff suppressed because it is too large Load Diff
@@ -13,7 +13,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""This module defines the model_config format.
"""This module defines the TensorRT-LLM checkpoint model_config format.
This format can be converted from huggingface, megatron or modelopt-quantized model.
And we will build tensorrt_llm engine from the context saved with this format.
@@ -26,38 +26,44 @@ import torch
from modelopt.torch.quantization.qtensor import NVFP4QTensor
QUANTIZATION_NONE = None
QUANTIZATION_FP8 = "fp8"
QUANTIZATION_INT8_SQ = "int8_sq"
QUANTIZATION_INT8_WO = "int8_wo"
QUANTIZATION_INT4_AWQ = "int4_awq"
QUANTIZATION_W4A8_AWQ = "w4a8_awq"
QUANTIZATION_NVFP4 = "nvfp4"
QUANTIZATION_NVFP4_SVDQUANT = "nvfp4_svdquant"
QUANTIZATION_W4A8_NVFP4_FP8 = "w4a8_nvfp4_fp8"
QUANTIZATION_MXFP4 = "mxfp4"
QUANTIZATION_MXFP8 = "mxfp8"
QUANTIZATION_W4A8_MXFP4_FP8 = "w4a8_mxfp4_fp8"
QUANTIZATION_W4A16_NVFP4 = "w4a16_nvfp4"
QUANTIZATION_NVFP4_AWQ = "nvfp4_awq"
QUANTIZATION_FP8_PB_REAL = "fp8_pb_real"
QUANTIZATION_FP8_PB_WO = "fp8_pb_wo"
QUANTIZATION_FP8_PC_PT = "fp8_pc_pt"
from ..quant_format import (
QUANTIZATION_INT4_AWQ,
QUANTIZATION_NONE,
QUANTIZATION_NVFP4,
QUANTIZATION_NVFP4_AWQ,
QUANTIZATION_NVFP4_SVDQUANT,
QUANTIZATION_W4A8_AWQ,
)
# Formats whose scales are purely per-module, so export never merges them across the q/k/v
# and gate/up groups that share an input. Every other format unifies input_amax (and, for
# NVFP4, weight_scale_2) across such a group, which only a whole-model forward can discover.
FUSION_FREE_FORMATS = frozenset({QUANTIZATION_FP8, QUANTIZATION_NONE, QUANTIZATION_FP8_PB_REAL})
__all__ = [
"LAYERNORM_DEFAULT",
"LAYERNORM_RMS",
"LINEAR_COLUMN",
"LINEAR_GROUP",
"LINEAR_ROW",
"AttentionConfig",
"ConvConfig",
"DecoderLayerConfig",
"EmbeddingConfig",
"ExpertConfig",
"LayernormConfig",
"LinearActConfig",
"LinearConfig",
"MLPConfig",
"MOEConfig",
"MedusaHeadConfig",
"ModelConfig",
"QKVConfig",
"RecurrentConfig",
"RelativeAttentionTableConfig",
"RgLruConfig",
]
KV_CACHE_FP8 = "FP8"
KV_CACHE_INT8 = "INT8"
KV_CACHE_NVFP4 = "NVFP4"
KV_CACHE_NVFP4_AFFINE = "NVFP4_AFFINE"
LINEAR_COLUMN = "column"
LINEAR_ROW = "row"
LINEAR_GROUP = "group"
# These need to be synced with torch.export.tensorrt_llm_type
# These need to be synced with modelopt.torch.export.trtllm.tensorrt_llm_type
LAYERNORM_DEFAULT = "LayerNorm"
LAYERNORM_RMS = "RmsNorm"
@@ -19,6 +19,7 @@ import copy
import json
import math
import tempfile
import warnings
from collections.abc import Iterator
from dataclasses import asdict
from pathlib import Path
@@ -32,6 +33,9 @@ from safetensors.torch import save_file
from modelopt.torch.utils import distributed as dist
from modelopt.torch.utils import import_plugin
from ..layer_utils import is_layernorm
from ..quant_format import QUANTIZATION_INT4_AWQ, QUANTIZATION_W4A8_AWQ
from ..quant_utils import get_quantization_format, process_layer_quant_config
from .layer_utils import (
build_conv_config,
build_decoder_config,
@@ -48,11 +52,10 @@ from .layer_utils import (
is_conv,
is_decoder_list,
is_embedding,
is_layernorm,
is_linear,
model_type_is_enc_dec,
)
from .model_config import QUANTIZATION_INT4_AWQ, QUANTIZATION_W4A8_AWQ, ModelConfig
from .model_config import ModelConfig
from .model_config_utils import (
merge_gate_fc,
merge_qkv,
@@ -67,7 +70,6 @@ from .postprocess import (
postprocess_tensors,
update_lm_head_quantization,
)
from .quant_utils import get_quantization_format, process_layer_quant_config
from .tensorrt_llm_utils import (
convert_to_tensorrt_llm_config,
is_tensorrt_llm_0_8_or_9,
@@ -84,6 +86,15 @@ with import_plugin("megatron", verbose=False):
__all__ = ["export_tensorrt_llm_checkpoint", "torch_to_tensorrt_llm_checkpoint"]
# Deprecated in 0.48.0, scheduled for removal in 0.49.0. The docs have carried a deprecation
# notice since 0.39.0 (2025-11-13), but this is the first release to warn at runtime and to say
# so in the changelog, so 0.48.0 starts the migration period the deprecation policy requires.
_DEPRECATION_MSG = (
"{name} and the TensorRT-LLM checkpoint format are deprecated as of 0.48.0 and will be "
"removed in 0.49.0. Use modelopt.torch.export.export_hf_checkpoint instead, which exports a "
"unified Hugging Face checkpoint deployable on TensorRT-LLM, vLLM and SGLang."
)
def torch_to_tensorrt_llm_checkpoint(
model: nn.Module,
@@ -95,6 +106,11 @@ def torch_to_tensorrt_llm_checkpoint(
) -> Iterator[tuple[dict[str, Any], dict[str, torch.Tensor], dict[str, Any]]]:
"""Converts the torch model to the TensorRT-LLM checkpoint per GPU rank.
.. deprecated:: 0.48.0
The TensorRT-LLM checkpoint format is deprecated and will be removed in 0.49.0. Use
:meth:`export_hf_checkpoint <modelopt.torch.export.unified_export_hf.export_hf_checkpoint>`
instead.
TensorRT-LLM checkpoint is the LLM model format that can be used by the TensorRT-LLM build API.
for the engine building process.
https://github.com/NVIDIA/TensorRT-LLM/blob/main/docs/source/architecture/checkpoint.md
@@ -120,6 +136,32 @@ def torch_to_tensorrt_llm_checkpoint(
per_layer_quantization: A dict that contains layer-wise quantization information for all quantized layers
for mixed_precision, empty dictionary otherwise.
"""
# Warn here rather than in the generator body: the body does not run until the first
# ``next()``, so a warning inside it would fire late (or never, if the caller never iterates).
warnings.warn(
_DEPRECATION_MSG.format(name="torch_to_tensorrt_llm_checkpoint"),
DeprecationWarning,
stacklevel=2,
)
return _torch_to_tensorrt_llm_checkpoint(
model=model,
decoder_type=decoder_type,
dtype=dtype,
inference_tensor_parallel=inference_tensor_parallel,
inference_pipeline_parallel=inference_pipeline_parallel,
workspace_path=workspace_path,
)
def _torch_to_tensorrt_llm_checkpoint(
model: nn.Module,
decoder_type: str,
dtype: torch.dtype | None = None,
inference_tensor_parallel: int = 0,
inference_pipeline_parallel: int = 1,
workspace_path: Path | str | None = None,
) -> Iterator[tuple[dict[str, Any], dict[str, torch.Tensor], dict[str, Any]]]:
"""Generator behind :func:`torch_to_tensorrt_llm_checkpoint`; see it for the contract."""
if dtype is None:
dtype = get_dtype(model)
@@ -447,6 +489,12 @@ def export_tensorrt_llm_checkpoint(
):
"""Exports the torch model to the TensorRT-LLM checkpoint and save to the export_dir.
.. deprecated:: 0.48.0
The TensorRT-LLM checkpoint format is deprecated and will be removed in 0.49.0. Use
:meth:`export_hf_checkpoint <modelopt.torch.export.unified_export_hf.export_hf_checkpoint>`
instead, which exports a unified Hugging Face checkpoint deployable on TensorRT-LLM,
vLLM and SGLang.
Args:
model: the torch model.
decoder_type: the type of the decoder, e.g. gpt, gptj, llama.
@@ -470,6 +518,11 @@ def export_tensorrt_llm_checkpoint(
https://github.com/NVIDIA/TensorRT-LLM/blob/main/tensorrt_llm/models/modeling_utils.py.
* ``.safetensors``: The file for the list of weights as safetensors. Unique for each rank.
"""
warnings.warn(
_DEPRECATION_MSG.format(name="export_tensorrt_llm_checkpoint"),
DeprecationWarning,
stacklevel=2,
)
export_dir = Path(export_dir)
export_root = export_dir
export_dir.mkdir(parents=True, exist_ok=True)
@@ -484,7 +537,7 @@ def export_tensorrt_llm_checkpoint(
tensorrt_llm_config,
weights,
quant_config,
) in torch_to_tensorrt_llm_checkpoint(
) in _torch_to_tensorrt_llm_checkpoint(
model=model,
decoder_type=decoder_type,
dtype=dtype,
@@ -23,10 +23,9 @@ from typing import Union, get_args, get_origin
import numpy as np
import torch
from ..quant_format import QUANTIZATION_FP8_PC_PT, QUANTIZATION_INT4_AWQ, QUANTIZATION_W4A8_AWQ
from ..quant_utils import to_quantized_weight
from .model_config import (
QUANTIZATION_FP8_PC_PT,
QUANTIZATION_INT4_AWQ,
QUANTIZATION_W4A8_AWQ,
DecoderLayerConfig,
LayernormConfig,
LinearConfig,
@@ -35,7 +34,6 @@ from .model_config import (
MOEConfig,
QKVConfig,
)
from .quant_utils import to_quantized_weight
# numpy doesn't know bfloat16, define abstract binary type instead
np_bfloat16 = np.dtype("V2", metadata={"dtype": "bfloat16"})
@@ -28,14 +28,12 @@ from modelopt.torch.quantization.nn.modules.quant_linear import QuantLinear
from modelopt.torch.quantization.qtensor import NVFP4QTensor
from modelopt.torch.utils import distributed as dist
from ..quant_format import QUANTIZATION_NVFP4, QUANTIZATION_NVFP4_AWQ, QUANTIZATION_NVFP4_SVDQUANT
from .distribute import get_configs_parallel, get_tensors_parallel
from .model_config import (
LINEAR_COLUMN,
LINEAR_GROUP,
LINEAR_ROW,
QUANTIZATION_NVFP4,
QUANTIZATION_NVFP4_AWQ,
QUANTIZATION_NVFP4_SVDQUANT,
ConvConfig,
EmbeddingConfig,
ExpertConfig,
+107
View File
@@ -0,0 +1,107 @@
# 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.
"""Quantization helpers used only by the TensorRT-LLM checkpoint export path."""
import torch
from modelopt.torch.quantization.qtensor import NVFP4QTensor
from ..quant_format import QUANTIZATION_NONE, QUANTIZATION_NVFP4_AWQ, QUANTIZATION_NVFP4_SVDQUANT
def get_scaling_factor_from_weight(weight, group_size) -> torch.tensor:
"""Calculate the weight scaling factor for a given group size."""
[n, k] = weight.shape
if group_size != 0:
# int4_awq
if k % group_size != 0:
raise NotImplementedError(
"Weight shape is not divisible for block size for block quantization."
)
weight = weight.reshape(n, k // group_size, group_size)
maxbound = 7.0
else:
# int8_sq
maxbound = 127.0
amax = weight.abs().max(dim=-1)[0].float()
weights_scaling_factor = amax / maxbound
# Let's filter the zeros in the scaling factor if the weights are zero
# to avoid the divided-by-zero error..
weights_scaling_factor[weights_scaling_factor == 0] = 1.0
return weights_scaling_factor
def resmooth_and_get_scale(
merged_weights: torch.Tensor,
pre_quant_scales: list[torch.Tensor],
ranks: int,
group_size: int,
new_pre_quant_scale: torch.Tensor | None = None,
quantization: str | None = QUANTIZATION_NONE,
):
"""Resmooths weights from a single or multiple ranks and get scaling factors and amax.
Args:
merged_weights: Merged weights from ranks.
pre_quant_scales: List of pre-quantization scales for each rank.
ranks: Number of ranks.
group_size: Group size of the quantization block.
new_pre_quant_scale (optional): If not provided, weights will be resmoothed using
the average of pre_quant_scales.
Returns:
weights: Resmoothed weights.
weight_scaling_factors: Resmoothed scaling factors.
avg_pre_quant_scale: Calculated average of the quantization scale.
"""
if new_pre_quant_scale is None:
new_pre_quant_scale = torch.stack(pre_quant_scales).mean(dim=0)
assert len(pre_quant_scales) > 0 and new_pre_quant_scale.numel() == merged_weights.shape[1], (
"Shape of pre_quant_scales and weights do not match."
)
weights = torch.chunk(merged_weights, ranks, dim=0)
scales = []
new_weights = []
for i, p_scaling_factor in enumerate(pre_quant_scales):
# De smooth & Re smooth
weight = (
weights[i]
* p_scaling_factor.type(weights[i].dtype)
/ new_pre_quant_scale.type(weights[i].dtype)
)
new_weights.append(weight)
# If NVFP4_AWQ then we view the scales as uint8 to allow for cat later
if quantization in [QUANTIZATION_NVFP4_AWQ, QUANTIZATION_NVFP4_SVDQUANT]:
scale, _ = NVFP4QTensor.get_weights_scaling_factor(weight, group_size).view(torch.uint8)
else:
scale = get_scaling_factor_from_weight(weight, group_size)
scales.append(scale)
resmoothed_scales = torch.cat(scales, dim=0)
return (
torch.cat(new_weights, dim=0),
resmoothed_scales.view(torch.float8_e4m3fn)
if quantization in [QUANTIZATION_NVFP4_AWQ, QUANTIZATION_NVFP4_SVDQUANT]
else resmoothed_scales, # if NVFP4_AWQ we view the scales back as float8_e4m3fn after cat
new_pre_quant_scale,
)
@@ -31,10 +31,10 @@ if TYPE_CHECKING:
from modelopt import __version__
from ..quant_format import QUANTIZATION_NONE
from .model_config import (
LAYERNORM_DEFAULT,
LAYERNORM_RMS,
QUANTIZATION_NONE,
DecoderLayerConfig,
MLPConfig,
ModelConfig,
+8 -8
View File
@@ -88,7 +88,14 @@ from .layer_utils import (
is_quantlinear,
sync_moe_gate_up_amax,
)
from .model_config import (
from .model_utils import TiedWeightMap, get_language_model_from_vl, is_multimodal_model
from .plugins import SpeculativeDecodingExporter, has_spec_opt, sanitize_hf_config_for_deployment
from .quant_aware_conversion import (
build_reverse_name_mapper,
revert_quant_config_names,
revert_weight_conversion_quant_aware,
)
from .quant_format import (
FUSION_FREE_FORMATS,
QUANTIZATION_FP8,
QUANTIZATION_FP8_PB_REAL,
@@ -102,13 +109,6 @@ from .model_config import (
QUANTIZATION_W4A8_NVFP4_FP8,
QUANTIZATION_W4A16_NVFP4,
)
from .model_utils import TiedWeightMap, get_language_model_from_vl, is_multimodal_model
from .plugins import SpeculativeDecodingExporter, has_spec_opt, sanitize_hf_config_for_deployment
from .quant_aware_conversion import (
build_reverse_name_mapper,
revert_quant_config_names,
revert_weight_conversion_quant_aware,
)
from .quant_utils import (
fuse_prequant_layernorm,
fuse_prequant_to_linear,
@@ -39,16 +39,6 @@ from modelopt.torch.quantization.nn.modules.tensor_quantizer import GroupedQuant
from modelopt.torch.utils import import_plugin, warn_rank_0
from .convert_hf_config import convert_hf_quant_config_format
from .model_config import (
KV_CACHE_FP8,
KV_CACHE_NVFP4,
QUANTIZATION_FP8,
QUANTIZATION_FP8_PB_REAL,
QUANTIZATION_FP8_PB_WO,
QUANTIZATION_NONE,
QUANTIZATION_NVFP4,
QUANTIZATION_W4A16_NVFP4,
)
from .plugins.hf_checkpoint_utils import (
copy_hf_ckpt_remote_code,
copy_non_safetensor_files_from_ckpt,
@@ -65,6 +55,16 @@ from .plugins.mcore_custom import (
save_safetensors_by_layer_index,
)
from .plugins.megatron_importer import GPTModelImporter, _get_mamba_conv1d
from .quant_format import (
KV_CACHE_FP8,
KV_CACHE_NVFP4,
QUANTIZATION_FP8,
QUANTIZATION_FP8_PB_REAL,
QUANTIZATION_FP8_PB_WO,
QUANTIZATION_NONE,
QUANTIZATION_NVFP4,
QUANTIZATION_W4A16_NVFP4,
)
from .quant_utils import (
get_activation_scaling_factor,
get_kv_cache_dtype,
+1 -38
View File
@@ -35,7 +35,7 @@ from _test_utils.torch.export.utils import (
from _test_utils.torch.transformers_models import get_tiny_qwen3_moe
import modelopt.torch.quantization as mtq
from modelopt.torch.export.model_config import (
from modelopt.torch.export.quant_format import (
KV_CACHE_FP8,
KV_CACHE_INT8,
QUANTIZATION_FP8,
@@ -52,7 +52,6 @@ from modelopt.torch.export.quant_utils import (
get_quant_config,
get_quantization_format,
get_scaling_factor,
get_scaling_factor_from_weight,
get_weight_block_size,
postprocess_state_dict,
process_layer_quant_config,
@@ -168,42 +167,6 @@ def test_all_items_same(item_list, expected):
assert generated == expected
@pytest.mark.parametrize(
("weight", "group_size", "expected"),
[
(
torch.tensor([[0.0, 0.35, 0.28, 7.0], [0.49, 0.84, -0.77, 0.07]]),
2,
torch.tensor([[0.05, 1.0], [0.12, 0.11]]),
), # group_size != 0 and divides weight.shape[1]
(
torch.tensor([[0.127, 0.0, 1.27, -12.7], [0.0, 127.0, 0.254, 2.54]]),
0,
torch.tensor([0.1, 1.0]),
), # group_size = 0
(
torch.tensor([[0.0, 0.0, 0.0, 0.0], [0.0, -0.127, 0.254, 2.54]]),
0,
torch.tensor([1.0, 0.02]),
), # zero replaced with 1.0
(
torch.tensor([[0.0, 0.84, -0.77, 0.07], [0.0, 0.0, 0.0, 0.0]]),
2,
torch.tensor([[0.12, 0.11], [1.0, 1.0]]),
), # zero replaced with 1.0
],
)
def test_get_scaling_factor_from_weight(weight, group_size, expected):
scaling_factor = get_scaling_factor_from_weight(weight, group_size)
# Check if shapes match
if group_size != 0:
assert list(scaling_factor.shape) == [weight.shape[0], weight.shape[1] // group_size]
else:
assert list(scaling_factor.shape) == [weight.shape[0]]
assert torch.allclose(scaling_factor, expected, rtol=0.0, atol=0.0)
@pytest.mark.parametrize(
("state_dict", "quantization", "maxbound", "expected_state_dict"),
[
@@ -26,13 +26,16 @@ from _test_utils.torch.export.utils import (
import modelopt.torch.export.unified_export_megatron as unified_export_megatron
import modelopt.torch.quantization as mtq
from modelopt.torch.export.layer_utils import get_quantization_format
from modelopt.torch.export.model_config import (
from modelopt.torch.export.quant_format import (
QUANTIZATION_FP8,
QUANTIZATION_NVFP4,
QUANTIZATION_W4A8_AWQ,
)
from modelopt.torch.export.quant_utils import get_kv_cache_scaling_factor, get_quant_config
from modelopt.torch.export.quant_utils import (
get_kv_cache_scaling_factor,
get_quant_config,
get_quantization_format,
)
from modelopt.torch.quantization.nn import NVFP4StaticQuantizer
@@ -36,8 +36,8 @@ from _test_utils.torch.quantization.tied_modules import (
)
import modelopt.torch.quantization as mtq
from modelopt.torch.export.model_config import KV_CACHE_FP8
from modelopt.torch.export.model_utils import TiedWeightMap
from modelopt.torch.export.quant_format import KV_CACHE_FP8
from modelopt.torch.export.quant_utils import _postprocess_single_tensor
from modelopt.torch.export.unified_export_hf import _export_quantized_weight
from modelopt.torch.export.unified_export_hf_streaming import (
@@ -0,0 +1,128 @@
# 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.
"""Unit tests for :mod:`modelopt.torch.export.trtllm.quant_utils`.
These helpers are reachable only from the TensorRT-LLM checkpoint export path
(``trtllm/postprocess.py``), so they are tested alongside it rather than with the
backend-agnostic helpers in ``modelopt.torch.export.quant_utils``.
Test modules in this directory carry a ``trtllm_`` prefix because pytest runs without
``__init__.py`` here and derives the module name from the bare filename, so a plain
``test_quant_utils.py`` would collide with the one a directory up.
"""
import inspect
import warnings
import pytest
import torch
import torch.nn as nn
import modelopt.torch.export as mte
from modelopt.torch.export.trtllm import (
export_tensorrt_llm_checkpoint,
model_config_export,
torch_to_tensorrt_llm_checkpoint,
)
from modelopt.torch.export.trtllm.quant_utils import get_scaling_factor_from_weight
@pytest.mark.parametrize(
("weight", "group_size", "expected"),
[
(
torch.tensor([[0.0, 0.35, 0.28, 7.0], [0.49, 0.84, -0.77, 0.07]]),
2,
torch.tensor([[0.05, 1.0], [0.12, 0.11]]),
), # group_size != 0 and divides weight.shape[1]
(
torch.tensor([[0.127, 0.0, 1.27, -12.7], [0.0, 127.0, 0.254, 2.54]]),
0,
torch.tensor([0.1, 1.0]),
), # group_size = 0
(
torch.tensor([[0.0, 0.0, 0.0, 0.0], [0.0, -0.127, 0.254, 2.54]]),
0,
torch.tensor([1.0, 0.02]),
), # zero replaced with 1.0
(
torch.tensor([[0.0, 0.84, -0.77, 0.07], [0.0, 0.0, 0.0, 0.0]]),
2,
torch.tensor([[0.12, 0.11], [1.0, 1.0]]),
), # zero replaced with 1.0
],
)
def test_get_scaling_factor_from_weight(weight, group_size, expected):
scaling_factor = get_scaling_factor_from_weight(weight, group_size)
# Check if shapes match
if group_size != 0:
assert list(scaling_factor.shape) == [weight.shape[0], weight.shape[1] // group_size]
else:
assert list(scaling_factor.shape) == [weight.shape[0]]
assert torch.allclose(scaling_factor, expected, rtol=0.0, atol=0.0)
def test_old_top_level_import_still_works():
"""The pre-0.48 import path stays importable for the migration period.
A DeprecationWarning is only useful if callers can still reach the code it warns about,
so removing this re-export before 0.49.0 would silently skip the migration window.
"""
assert mte.export_tensorrt_llm_checkpoint is export_tensorrt_llm_checkpoint
assert mte.torch_to_tensorrt_llm_checkpoint is torch_to_tensorrt_llm_checkpoint
def test_torch_to_tensorrt_llm_checkpoint_warns_at_call_time():
"""The warning must fire on call, not on first ``next()``.
``torch_to_tensorrt_llm_checkpoint`` hands back a generator. A bare ``warnings.warn`` in a
generator body does not run until the first item is pulled, so a caller that builds the
generator and abandons it would never be warned. Hence the public name is a plain function
wrapping a private generator.
"""
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
generator = torch_to_tensorrt_llm_checkpoint(nn.Linear(4, 4), "llama")
assert inspect.isgenerator(generator)
deprecations = [w for w in caught if issubclass(w.category, DeprecationWarning)]
assert len(deprecations) == 1, "expected exactly one DeprecationWarning before iteration"
assert "torch_to_tensorrt_llm_checkpoint" in str(deprecations[0].message)
def test_export_tensorrt_llm_checkpoint_warns_exactly_once(tmp_path, monkeypatch):
"""One user call yields one warning, not two.
``export_tensorrt_llm_checkpoint`` drives the same generator, so it must call the private
``_torch_to_tensorrt_llm_checkpoint``; going through the public wrapper would emit a second,
redundant warning.
The generator is stubbed to yield nothing so the call completes normally. Letting a real
conversion fail and swallowing the exception would also pass, but it would pass for any
failure after the warning, which is not what this test is about.
"""
monkeypatch.setattr(
model_config_export, "_torch_to_tensorrt_llm_checkpoint", lambda **kwargs: iter(())
)
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
export_tensorrt_llm_checkpoint(nn.Linear(4, 4), "llama", export_dir=tmp_path)
deprecations = [w for w in caught if issubclass(w.category, DeprecationWarning)]
assert len(deprecations) == 1, f"expected 1 DeprecationWarning, got {len(deprecations)}"
assert "export_tensorrt_llm_checkpoint" in str(deprecations[0].message)