Files
Model-Optimizer/modelopt/torch/export/trtllm/distribute.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

306 lines
11 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.
"""torch.distribute utils."""
import json
from contextlib import contextmanager
from io import BytesIO
from multiprocessing.shared_memory import SharedMemory
from pathlib import Path
from typing import Any
import torch
from modelopt.torch.utils import distributed as dist
from modelopt.torch.utils import safe_load
from .model_config_utils import (
model_config_from_dict,
model_config_to_dict,
restore_model_config,
split_config_and_weights,
)
class NFSWorkspace:
"""A shared workspace implementation using Network File Storage (NFS).
NOTE: all read/write/modifition to the NFS dir do not involve any collective
communication nor barrier. It is users' responsibility to synchronize
all ranks (local and remove processes).
This implementation uses `torch.save` and `safe_load` (`torch.load(weights_only=True)`) for serialization.
Args:
workspace_path: the path to the NFS directory for postprocess cross rank communication.
If not provided, SharedMemory will be used instead.
"""
def __init__(self, workspace_path: Path | str | None = None):
"""Create the NFS work dir and clean up existing existing state files."""
self.path = Path("") if workspace_path is None else Path(workspace_path)
self._is_initialized = workspace_path is not None
self.rank = dist.rank()
if self.is_initialized:
if self.rank == 0:
self.path.mkdir(parents=True, exist_ok=True)
self.state_path = self._get_state_path(self.rank)
self._clean_up()
@property
def is_initialized(self):
"""Whether the workspace is initialized."""
return self._is_initialized
def write_configs_and_weights(self, config_json: dict[str, Any], weights: dict[str, Any]):
"""All ranks write the state file to the shared NFS dir.
Args:
config_json: model or module config in json
weights: module weights in torch's state_dict format
"""
if not self.is_initialized:
raise ValueError("NFSWorkspace is not initialized!")
self._clean_up()
torch.save({"config": config_json, "weight": weights}, self.state_path)
def read_configs_and_weights_from_rank(
self, target_rank: int
) -> tuple[dict[str, Any] | None, dict[str, Any] | None]:
"""All ranks read the target_rank state file.
Args:
target_rank: the target rank
Returns:
the model/module config and the weights
"""
if not self.is_initialized:
raise ValueError("NFSWorkspace is not initialized!")
state_path = self._get_state_path(target_rank)
if state_path.exists():
state = safe_load(state_path, map_location="cpu")
return state["config"], state["weight"]
else:
return None, None
def _get_state_path(self, target_rank: int) -> Path:
"""Return the state file name of a particular rank.
Args:
target_rank: the target rank
Returns:
the state file path of the target rank
"""
if not self.is_initialized:
raise ValueError("NFSWorkspace is not initialized!")
return self.path.joinpath(f"rank_{target_rank}_state.pth")
def _clean_up(self):
"""Remove existing state files."""
if not self.is_initialized:
raise ValueError("NFSWorkspace is not initialized!")
self.state_path.unlink(missing_ok=True)
@contextmanager
def get_tensors_parallel(tensor: torch.Tensor, ranks: list[int], group=None):
"""Gathers the tensors across distributed processes using shm.
Args:
tensor: the tensor that each rank want to pass to the first rank.
The tensors across the ranks need to have the same size.
ranks: the list of the ranks
group: the barrier sync group.
Yields:
the first rank in the ranks has the full access of the tensors across all the ranks.
the other ranks returns an empty list
The shm will be destroyed after consumption.
"""
assert tensor is not None
assert len(ranks) > 1
local_rank = dist.rank()
shm_writer = None
shm_readers = []
tensor = tensor.cpu()
is_merged_rank = local_rank == ranks[0]
# Create shm and copy the tensor to the shm if not the merged rank.
# Assume each tensor need up to 2KB additional space for metadata.
if not is_merged_rank:
shm_writer = SharedMemory(name=f"rank_{local_rank}", create=True, size=tensor.nbytes + 2048)
torch.save(tensor, shm_writer._mmap) # type: ignore[attr-defined]
# All ranks wait for this to complete.
dist.barrier(group)
tensors = []
# The merged rank gather the tensor from the other ranks (including itself).
if is_merged_rank:
for rank in ranks:
if rank == ranks[0]:
tensors.append(tensor)
else:
shm = SharedMemory(name=f"rank_{rank}", create=False)
assert shm.buf is not None
shared_tensor = torch.load(BytesIO(shm.buf))
tensors.append(shared_tensor)
shm_readers.append(shm)
try:
# Send the tensor list to the consumer.
# The merged rank will get a valid tensor while the other ranks an empty tensor.
yield tensors
finally:
# Reader closes the shms.
if shm_readers:
for shm in shm_readers:
shm.close()
# All ranks wait for the reader to close the shms.
dist.barrier(group)
# Writer frees the shm resource.
if shm_writer is not None:
shm_writer.close()
shm_writer.unlink()
@contextmanager
def get_configs_parallel(config, ranks: list[int], group, workspace_path: Path | str | None = None):
"""Gathers the layer config across distributed processes using shm or NFS.
Args:
config: the config (nullable) that each rank want to pass to the first rank.
ranks: the list of the ranks
group: the barrier sync group.
workspace_path: the path to the NFS directory for postprocess cross rank communication.
Yields:
the first rank in the ranks has the full access of the configs across all the ranks.
the other ranks returns an empty list
When workspace_path is provided, an NFSWorkspace object is created to perform communication
across ranks. Otherwise, `SharedMemory` is used for local multi-process communication.
The shm will be destroyed after consumption.
"""
assert len(ranks) > 1
local_rank = dist.rank()
shm_writer = None
shm_readers = []
nfs_workspace = NFSWorkspace(workspace_path)
is_merged_rank = local_rank == ranks[0]
def _get_weights_nbytes(weights_dict: dict[str, torch.Tensor]):
total_nbytes = 0
for k, v in weights_dict.items():
# Assume each tensor need up to 2KB additional space for metadata.
# In reality this should be much smaller.
total_nbytes = total_nbytes + len(k) + v.nbytes + 2048
return total_nbytes
# Create shm and copy the serialized config to the shm if not the merged rank.
if not is_merged_rank:
if config is not None:
config_dict = model_config_to_dict(config)
# Add additional config type name to the dict so we can later pick the right config type.
config_dict["__name__"] = str(type(config).__name__)
weights = {}
split_config_and_weights(config_dict, weights)
config_json = json.dumps(config_dict)
if nfs_workspace.is_initialized:
# All ranks except for the master merge rank write to the NFS dir.
nfs_workspace.write_configs_and_weights(config_dict, weights)
else:
# SHM data structure: 8B json size, serialized json bytes and the weights dict.
shm_writer = SharedMemory(
name=f"rank_{local_rank}_config",
create=True,
size=(8 + len(config_json) + _get_weights_nbytes(weights)),
)
assert shm_writer.buf is not None
# Write json length to the shm
shm_writer.buf[:8] = len(config_json).to_bytes(8, "little")
# Write json to the shm
shm_writer.buf[8 : len(config_json) + 8] = config_json.encode()
# Write np tensors to the shm.
shm_writer._mmap.seek(len(config_json) + 8) # type: ignore[attr-defined]
torch.save(weights, shm_writer._mmap) # type: ignore[attr-defined]
else:
# If the config is None, we just store the empty 0.
shm_writer = SharedMemory(
name=f"rank_{local_rank}_config",
create=True,
size=8,
)
assert shm_writer.buf is not None
shm_writer.buf[:8] = (0).to_bytes(8, "little")
# All ranks wait for this to complete.
dist.barrier(group)
configs = []
if is_merged_rank:
for rank in ranks:
if rank == ranks[0]:
configs.append(config)
elif nfs_workspace.is_initialized:
# The master merge rank read other configs from the NFS dir.
config_dict, weights = nfs_workspace.read_configs_and_weights_from_rank(rank)
if config_dict is not None:
restore_model_config(config_dict, weights)
config = model_config_from_dict(config_dict)
configs.append(config)
else:
shm = SharedMemory(name=f"rank_{rank}_config", create=False)
assert shm.buf is not None
len_json = int.from_bytes(shm.buf[:8], "little")
if len_json != 0:
config_dict = json.loads(shm.buf[8 : 8 + len_json].tobytes().decode())
weights = torch.load(BytesIO(shm.buf[8 + len_json :]))
restore_model_config(config_dict, weights)
config = model_config_from_dict(config_dict)
configs.append(config)
shm_readers.append(shm)
try:
# Send the config list to the consumer.
# The merged rank will get a valid config list while the other ranks an empty list.
yield configs
finally:
# Reader closes the shms.
if shm_readers:
for shm in shm_readers:
shm.close()
# All ranks wait for the reader to close the shms.
dist.barrier(group)
# Writer frees the shm resource.
if shm_writer is not None:
shm_writer.close()
shm_writer.unlink()