Fix 2-GPU test_model_load_utils hang; test import and fixture cleanup (#2079)

### What does this PR do?

Type of change: Bug fix + new tests (test-suite only; changes are
confined to `tests/`)

- **Fix `test_parallel_load_and_export` hang on 2 GPUs.** The temp paths
were built from `os.getpid()` *inside* the workers, so each rank got a
different `ckpt_dir`; rank 1 then took `_resolve_checkpoint_dir`'s Hub
branch and blocked on an extra barrier while rank 0 ran the loader's
broadcasts. The checkpoint is now built once in the parent under
`tmp_path` and passed in, so all ranks agree on the path.
- **Imports moved to module top** across `tests/gpu*` and
`tests/examples`; function-local imports kept only where guarded
(`importorskip`/`try`), where the import *is* the test (JIT compile), or
where it must follow `sys.path` setup.
- **Reuse `_test_utils` instead of local copies:** added
`get_tiny_mixtral`; deduped `assert_nodes_are_quantized` (5 copies), the
accelerate-offload/layerwise config helpers, `make_quant_attention`,
`get_dflash_config`, the NVFP4 amax assertions, and 3 copies of the
`tiny_wan22_path` fixture.
- **Shared model-dir fixtures assert they were not modified**
(`assert_unmodified_tree`): a file manifest is compared on teardown, so
a test that writes into a session-scoped fixture directory fails instead
of silently changing what later tests see.
- **Dropped dependency guards the CI env already guarantees**
(diffusers, tensorrt_llm in `gpu_trtllm`, transformer_engine in
`gpu_megatron`, transformers in examples) so a missing dep fails loudly
instead of skipping.
- **`test_heterogenous_sharded_state_dict` is skipped on Blackwell**
(sm_120), matching the existing marker and its TE/CUDA-13 rationale —
same tracking issue as #1901.

### Testing

Local, 2x RTX 6000 Ada: `tests/gpu/torch/utils/test_model_load_utils.py`
passes on 2 GPUs and on 1 GPU (previously hung on 2). Also ran the
touched files in `tests/unit` (608 passed) and `tests/gpu` (~340
passed).

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

- Is this change backward compatible?: ✅
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: N/A — test-only PR
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
N/A
- Did you get Claude approval on this PR?: ❌ — not yet run


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

* **Tests**
* Added reusable validation utilities for generated files, quantization
behavior, attention modules, model fixtures, offloading, and speculative
decoding.
* Expanded coverage for tiny Wan, Mixtral, Llama, and related model
scenarios.
* Consolidated duplicated setup and assertions across ONNX, GPU,
quantization, export, and sparsity tests.
* Improved fixture integrity checks, NVFP4 validation, and handling of
identity inputs.
* Reduced unnecessary dependency-based skips and isolated known
platform-specific flakiness.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
This commit is contained in:
Keval Morabia
2026-08-05 21:00:42 +05:30
committed by GitHub
parent 7afbfbc85c
commit 19e0121d0a
42 changed files with 530 additions and 592 deletions
+57
View File
@@ -0,0 +1,57 @@
# 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.
"""Filesystem helpers for tests."""
from collections.abc import Iterator
from contextlib import contextmanager
from pathlib import Path
def _manifest(root: Path) -> dict[str, tuple[int, int]]:
"""``relative path -> (size, mtime_ns)`` for every file below ``root``."""
return {
str(p.relative_to(root)): (p.stat().st_size, p.stat().st_mtime_ns)
for p in root.rglob("*")
if p.is_file() and not p.is_symlink()
}
@contextmanager
def assert_unmodified_tree(path: Path | str) -> Iterator[Path]:
"""Fail if anything under ``path`` is added, removed, or rewritten inside the ``with``.
For session/module-scoped model-directory fixtures: a test that writes into a shared
directory silently changes what every later test sees. Comparing a file manifest on
teardown catches that. ``chmod``-ing the tree read-only would report at the write rather
than at teardown, but it only works for an unprivileged user -- root has
``CAP_DAC_OVERRIDE`` and writes straight through the permission bits, and the CI
containers run as root.
"""
path = Path(path)
before = _manifest(path)
yield path
if not path.exists():
raise AssertionError(f"shared fixture directory {path} was deleted by a test")
after = _manifest(path)
added = sorted(after.keys() - before.keys())
removed = sorted(before.keys() - after.keys())
changed = sorted(k for k in before.keys() & after.keys() if before[k] != after[k])
if added or removed or changed:
raise AssertionError(
f"shared fixture directory {path} was modified by a test "
f"(added={added}, removed={removed}, changed={changed}); "
"copy it into the test's own tmp_path instead of writing into the shared tree"
)
+8 -1
View File
@@ -16,11 +16,18 @@
import onnx_graphsurgeon as gs
def assert_nodes_are_quantized(nodes):
def assert_nodes_are_quantized(nodes, *, ignore_identity_inputs: bool = False):
"""Assert every variable input of ``nodes`` is produced by a DequantizeLinear.
``ignore_identity_inputs`` skips inputs fed by an ``Identity`` node, for graphs where the
quantizer legitimately leaves such a passthrough in place (e.g. concat elimination).
"""
for node in nodes:
for inp_idx, inp in enumerate(node.inputs):
if isinstance(inp, gs.Variable):
producer = node.i(inp_idx)
if ignore_identity_inputs and producer and producer.op == "Identity":
continue
# Quantized path may include a Cast right after DQ
if producer and producer.op == "Cast":
producer = producer.i(0)
@@ -0,0 +1,37 @@
# 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.
"""Shared attention-quantization fixtures for the unit and gpu attention tests."""
import pytest
pytest.importorskip("transformers")
from transformers import LlamaConfig
from transformers.models.llama.modeling_llama import LlamaAttention
from modelopt.torch.quantization.plugins.huggingface import _QuantAttention
def make_quant_attention(hidden_size=128, num_q_heads=4, num_kv_heads=2):
"""A single ``_QuantAttention``-converted Llama attention layer, pinned to the sdpa impl."""
config = LlamaConfig(
hidden_size=hidden_size,
num_attention_heads=num_q_heads,
num_key_value_heads=num_kv_heads,
)
quant_attention = _QuantAttention.convert(LlamaAttention(config, layer_idx=0))
quant_attention.config._attn_implementation = "sdpa"
return quant_attention
@@ -0,0 +1,73 @@
# 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.
"""Shared helpers for accelerate-offloaded and layerwise-calibration quantization tests."""
import copy
import torch
from _test_utils.torch.transformers_models import create_tiny_llama_dir
from accelerate import init_empty_weights, load_checkpoint_and_dispatch
from transformers import AutoConfig, AutoModelForCausalLM
def make_tiny_llama_and_inputs(tmp_path, num_hidden_layers=3):
"""Tiny LLaMA checkpoint dir + its config + a GPU token batch sized for its vocab."""
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_hidden_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
inputs = torch.randint(0, config.vocab_size, (1, 4)).cuda()
return tiny_llama_dir, config, inputs
def make_cpu_offloaded_model(tmp_path, num_hidden_layers=3):
"""Tiny LLaMA with layer 0 offloaded to CPU via accelerate.
Returns ``(model, config, tiny_llama_dir, inputs)``; ``inputs`` is a GPU token batch
sized for the model's vocab.
"""
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_hidden_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
with init_empty_weights():
model = AutoModelForCausalLM.from_config(config)
device_map = {
n: 0
for n, m in model.named_modules()
if "layers" not in n or n.split("layers.")[-1].isdigit()
}
device_map["model.layers.0"] = "cpu"
model = load_checkpoint_and_dispatch(model, tiny_llama_dir, device_map=device_map)
inputs = torch.randint(0, config.vocab_size, (1, 4)).cuda()
return model, config, tiny_llama_dir, inputs
def make_layerwise_cfg(base_cfg):
"""Copy of ``base_cfg`` with ``layerwise=True`` set on its algorithm field."""
cfg = copy.deepcopy(base_cfg)
algo = cfg.get("algorithm", "max")
if isinstance(algo, str):
cfg["algorithm"] = {"method": algo, "layerwise": True}
else:
algo["layerwise"] = True
return cfg
def make_layerwise_checkpoint_cfg(base_cfg, checkpoint_dir):
"""``make_layerwise_cfg`` plus a ``layerwise_checkpoint_dir``."""
cfg = make_layerwise_cfg(base_cfg)
cfg["algorithm"]["layerwise_checkpoint_dir"] = checkpoint_dir
return cfg
@@ -15,6 +15,8 @@
import torch
from modelopt.torch.quantization.nn import TensorQuantizer
def quant(x, amax, num_bits=8, fake=False, narrow_range=True):
"""Quantize x using torch."""
@@ -32,3 +34,23 @@ def quant(x, amax, num_bits=8, fake=False, narrow_range=True):
def get_model_size(model):
return sum([p.element_size() * p.nelement() for p in model.parameters()])
def nvfp4_static_amax_dtypes(model):
"""Map of ``module name -> amax dtype`` for every NVFP4 static quantizer with an amax."""
return {
name: module.amax.dtype
for name, module in model.named_modules()
if isinstance(module, TensorQuantizer)
and module.is_nvfp4_static
and module.amax is not None
}
def assert_nvfp4_static_amaxes_fp32(amax_dtypes, model_dtype, label):
"""NVFP4 static amaxes must stay fp32 regardless of the model's own dtype."""
assert amax_dtypes, f"{label}: expected NVFP4 static amaxes for model dtype {model_dtype}"
assert all(amax_dtype == torch.float32 for amax_dtype in amax_dtypes.values()), (
f"{label}: expected all NVFP4 static amaxes to be fp32 for model dtype {model_dtype}, "
f"got {amax_dtypes}"
)
@@ -0,0 +1,42 @@
# 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.
"""Shared DFlash test config, used by the unit and gpu speculative-decoding tests."""
from copy import deepcopy
from modelopt.torch.speculative.config import DFLASH_DEFAULT_CFG
DFLASH_BLOCK_SIZE = 4
DFLASH_NUM_DRAFT_LAYERS = 2
def get_dflash_config(
block_size: int = DFLASH_BLOCK_SIZE,
num_layers: int = DFLASH_NUM_DRAFT_LAYERS,
offline: bool | None = None,
):
"""DFlash config sized for a tiny model: no torch.compile, token 0 as the mask token.
``offline`` is only written when set, so callers that don't care keep the default.
"""
config = deepcopy(DFLASH_DEFAULT_CFG["config"])
config["dflash_block_size"] = block_size
config["dflash_use_torch_compile"] = False
config["dflash_mask_token_id"] = 0 # use token 0 as mask for tiny model
config["dflash_architecture_config"] = {"num_hidden_layers": num_layers}
if offline is not None:
config["dflash_offline"] = offline
return config
+20 -4
View File
@@ -35,6 +35,7 @@ from transformers import (
GptOssConfig,
LlamaConfig,
LlamaForSequenceClassification,
MixtralConfig,
NemotronConfig,
PreTrainedModel,
Qwen3Config,
@@ -566,6 +567,25 @@ def create_tiny_gpt_oss_dir(
)
##### MIXTRAL #####
def get_tiny_mixtral(**config_kwargs) -> PreTrainedModel:
set_seed(SEED)
kwargs = {
"dtype": torch.bfloat16,
"hidden_size": 32,
"intermediate_size": 32,
"num_hidden_layers": 2,
"num_attention_heads": 4,
"num_key_value_heads": 2,
"num_local_experts": 4,
"num_experts_per_tok": 2,
"max_position_embeddings": 32,
"vocab_size": 32,
}
kwargs.update(config_kwargs)
return AutoModelForCausalLM.from_config(MixtralConfig(**kwargs))
##### LLAMA #####
def get_tiny_llama(**config_kwargs) -> PreTrainedModel:
set_seed(SEED)
@@ -671,10 +691,6 @@ def get_tiny_bert(**config_kwargs) -> PreTrainedModel:
return AutoModelForQuestionAnswering.from_config(BertConfig(**kwargs))
def create_tiny_bert_dir(tmp_path: Path | str, **config_kwargs) -> Path:
return _create_tiny_llm_dir(Path(tmp_path) / "tiny_bert", get_tiny_bert, **config_kwargs)
##### ViT (vision) #####
def get_tiny_vit(**config_kwargs) -> PreTrainedModel:
set_seed(SEED)
+28 -10
View File
@@ -20,6 +20,7 @@ from pathlib import Path
import pytest
import torch
import torch.distributed as dist
from _test_utils.fs_utils import assert_unmodified_tree
from _test_utils.torch.distributed.utils import init_process
import modelopt.torch.opt as mto
@@ -91,16 +92,7 @@ def pytest_collection_modifyitems(config, items):
item.add_marker(pytest.mark.timeout(_DEFAULT_TIMEOUT[group]))
@pytest.fixture
def tiny_tokenizer():
"""Real tiny HF tokenizer (vocab=128) shared across unit and gpu test lanes."""
# Lazy import: transformers_models.py runs ``pytest.importorskip("transformers")``
# at module load, which we don't want to trigger at conftest import time.
from _test_utils.torch.transformers_models import get_tiny_tokenizer
return get_tiny_tokenizer()
# General Fixtures #################################################################################
@pytest.fixture
def skip_on_windows():
if platform.system() == "Windows":
@@ -161,3 +153,29 @@ def enable_hf_checkpointing():
def project_root_path(request: pytest.FixtureRequest) -> Path:
"""Fixture providing the project root path for tests."""
return Path(request.config.rootpath)
# Transformers Models Fixtures #####################################################################
@pytest.fixture
def tiny_tokenizer():
"""Real tiny HF tokenizer (vocab=128) shared across unit and gpu test lanes."""
# Lazy import: transformers_models.py runs ``pytest.importorskip("transformers")``
# at module load, which we don't want to trigger at conftest import time.
from _test_utils.torch.transformers_models import get_tiny_tokenizer
return get_tiny_tokenizer()
@pytest.fixture(scope="session")
def tiny_wan22_path(tmp_path_factory):
"""Tiny Wan 2.2 pipeline dir, built once per session (the build is the expensive part).
Shared by the gpu sparse-attention tests and the diffusers example tests.
"""
# Lazy import for the same reason as ``tiny_tokenizer``: diffusers_models.py pulls in
# transformers at module load.
from _test_utils.torch.diffusers_models import create_tiny_wan22_pipeline_dir
pipeline_dir = create_tiny_wan22_pipeline_dir(tmp_path_factory.mktemp("tiny_wan22"))
with assert_unmodified_tree(pipeline_dir) as path:
yield str(path)
+21 -20
View File
@@ -15,6 +15,7 @@
import pytest
from _test_utils.fs_utils import assert_unmodified_tree
from _test_utils.torch.transformers_models import (
create_tiny_gpt_oss_dir,
create_tiny_llama_dir,
@@ -24,34 +25,34 @@ from _test_utils.torch.transformers_models import (
@pytest.fixture(scope="session")
def tiny_llama_path(tmp_path_factory):
return str(
create_tiny_llama_dir(
tmp_path_factory.mktemp("tiny_llama"),
with_tokenizer=True,
hidden_size=512,
intermediate_size=512,
)
model_dir = create_tiny_llama_dir(
tmp_path_factory.mktemp("tiny_llama"),
with_tokenizer=True,
hidden_size=512,
intermediate_size=512,
)
with assert_unmodified_tree(model_dir) as path:
yield str(path)
@pytest.fixture(scope="session")
def tiny_qwen3_path(tmp_path_factory):
return str(
create_tiny_qwen3_dir(
tmp_path_factory.mktemp("tiny_qwen3"),
with_tokenizer=True,
hidden_size=512,
intermediate_size=512,
)
model_dir = create_tiny_qwen3_dir(
tmp_path_factory.mktemp("tiny_qwen3"),
with_tokenizer=True,
hidden_size=512,
intermediate_size=512,
)
with assert_unmodified_tree(model_dir) as path:
yield str(path)
@pytest.fixture(scope="session")
def tiny_gpt_oss_path(tmp_path_factory):
return str(
create_tiny_gpt_oss_dir(
tmp_path_factory.mktemp("tiny_gpt_oss"),
with_tokenizer=True,
num_hidden_layers=2,
)
model_dir = create_tiny_gpt_oss_dir(
tmp_path_factory.mktemp("tiny_gpt_oss"),
with_tokenizer=True,
num_hidden_layers=2,
)
with assert_unmodified_tree(model_dir) as path:
yield str(path)
+4 -29
View File
@@ -14,37 +14,12 @@
# limitations under the License.
import pytest
@pytest.fixture(scope="session")
def tiny_wan22_path(tmp_path_factory):
"""Create a tiny Wan 2.2 (14B-style) pipeline and return its path.
Built once per session and shared across all tests that need it.
"""
try:
from _test_utils.torch.diffusers_models import create_tiny_wan22_pipeline_dir
except ImportError:
pytest.skip("Wan 2.2 diffusers models not available (requires diffusers with WanPipeline)")
tmp_path = tmp_path_factory.mktemp("wan22")
return str(create_tiny_wan22_pipeline_dir(tmp_path))
from _test_utils.fs_utils import assert_unmodified_tree
from _test_utils.torch.diffusers_models import create_tiny_qwen_image_pipeline_dir
@pytest.fixture(scope="session")
def tiny_qwen_image_path(tmp_path_factory):
"""Create a tiny Qwen-Image pipeline and return its path (built once per session).
Used by the diffusers Qwen export tests and the recipe-level DMD2 e2e
(``test_fastgen_recipe_e2e.py``). The pipeline is built fully offline by
``create_tiny_qwen_image_pipeline_dir`` (inline tiny Qwen2.5-VL text encoder +
local byte-level tokenizer); it skips only when the diffusers Qwen classes are
unavailable.
"""
try:
from _test_utils.torch.diffusers_models import create_tiny_qwen_image_pipeline_dir
except ImportError:
pytest.skip("Qwen-Image diffusers models not available")
tmp_path = tmp_path_factory.mktemp("qwen_image")
return str(create_tiny_qwen_image_pipeline_dir(tmp_path))
with assert_unmodified_tree(create_tiny_qwen_image_pipeline_dir(tmp_path)) as path:
yield str(path)
@@ -28,7 +28,11 @@ import math
import pytest
import torch
from _test_utils.examples.run_command import run_example_command
from _test_utils.torch.diffusers_models import create_tiny_wan22_pipeline_dir
from diffusers import AutoencoderKLWan, WanPipeline
import modelopt.torch.sparsity.attention_sparsity as mtsa
from modelopt.torch.export import export_hf_checkpoint
from modelopt.torch.sparsity.attention_sparsity.sparse_attention import SparseAttentionModule
EXAMPLE_PATH = "diffusers/sparsity"
@@ -51,12 +55,6 @@ _TINY_ARGS = [
]
@pytest.fixture(scope="session")
def tiny_wan22_path(tmp_path_factory):
"""Create a tiny Wan 2.2 pipeline saved to disk (session-scoped)."""
return str(create_tiny_wan22_pipeline_dir(tmp_path_factory.mktemp("tiny_wan22")))
def test_wan22_baseline(tiny_wan22_path, tmp_path):
"""Dense baseline — no sparsity, default diffusers attention backend."""
cmd = [
@@ -153,7 +151,7 @@ def test_wan22_export_sparse_checkpoint(tiny_wan22_path, tmp_path):
assert not (export_dir / "sparse.yaml").exists(), "Unexpected top-level sparse.yaml"
def test_wan22_calibrated_export(tmp_path):
def test_wan22_calibrated_export(tiny_wan22_path, tmp_path):
"""Inject calibration params via the Python API and verify the exported config.
Calibration can't succeed on tiny models via the Triton kernel (not enough
@@ -162,13 +160,7 @@ def test_wan22_calibrated_export(tmp_path):
(top-level ``threshold_scale_factor`` of the form ``a * exp(b * target_sparsity)``)
and that the dense (cross-attention) layers are recorded under ``ignore``.
"""
from diffusers import AutoencoderKLWan, WanPipeline
import modelopt.torch.sparsity.attention_sparsity as mtsa
from modelopt.torch.export import export_hf_checkpoint
from modelopt.torch.sparsity.attention_sparsity.sparse_attention import SparseAttentionModule
pipe_dir = create_tiny_wan22_pipeline_dir(tmp_path / "model")
pipe_dir = tiny_wan22_path
vae = AutoencoderKLWan.from_pretrained(pipe_dir, subfolder="vae", torch_dtype=torch.float32)
pipe = WanPipeline.from_pretrained(pipe_dir, vae=vae, torch_dtype=torch.bfloat16)
pipe.to("cuda")
@@ -14,6 +14,7 @@
# limitations under the License.
import json
import re
from pathlib import Path
from typing import NamedTuple
@@ -21,6 +22,7 @@ import pytest
from _test_utils.examples.models import FLUX_SCHNELL_PATH, SDXL_PATH
from _test_utils.examples.run_command import run_example_command
from _test_utils.torch.misc import minimum_sm
from safetensors import safe_open
class DiffuserHfExportModel(NamedTuple):
@@ -177,8 +179,6 @@ def _module_prefixes(keys: set[str], suffix: str) -> set[str]:
def _block_indices(prefixes: set[str]) -> set[int]:
"""transformer_blocks indices referenced by a set of module prefixes."""
import re
indices = set()
for prefix in prefixes:
match = re.search(r"transformer_blocks\.(\d+)\.", prefix)
@@ -227,8 +227,6 @@ _QWEN_SVDQUANT_SKIPPED_SUFFIXES = (
def test_qwen_image_hf_ckpt_export(
qwen_model: QwenHfExportModel, tiny_qwen_image_path: str, tmp_path: Path
) -> None:
from safetensors import safe_open
hf_ckpt_dir = qwen_model.quantize_and_export_hf(tiny_qwen_image_path, tmp_path)
assert hf_ckpt_dir.exists(), f"HF checkpoint directory was not created: {hf_ckpt_dir}"
@@ -21,8 +21,6 @@ bucket of configuration.json files. Worth exercising it explicitly.
"""
import pytest
pytest.importorskip("transformers") # utils.py imports AutoTokenizer at module load
from specdec_bench.utils import (
_SENSITIVE_KEY_ALLOWLIST,
_is_sensitive_key,
@@ -19,20 +19,11 @@ import sys
import onnx
import onnx_graphsurgeon as gs
from _test_utils.onnx.lib_test_models import build_conv_concat_model
from _test_utils.onnx.quantization.utils import assert_nodes_are_quantized
from modelopt.onnx.quantization.quantize import quantize
def assert_nodes_are_quantized(nodes):
for node in nodes:
for inp_idx, inp in enumerate(node.inputs):
if isinstance(inp, gs.Variable) and node.i(inp_idx).op != "Identity":
assert node.i(inp_idx).op == "DequantizeLinear", (
f"Input '{inp.name}' of node '{node.name}' is not quantized but should be!"
)
return True
def _check_concat_qdq_status(onnx_path, quantize_mode):
# Quantize the input model
quantize(onnx_path, quantize_mode=quantize_mode, passes="concat_elimination")
@@ -48,7 +39,7 @@ def _check_concat_qdq_status(onnx_path, quantize_mode):
# Check that all Conv nodes are quantized
conv_nodes = [n for n in graph.nodes if n.op == "Conv"]
assert assert_nodes_are_quantized(conv_nodes)
assert assert_nodes_are_quantized(conv_nodes, ignore_identity_inputs=True)
check_num = 0
for node in graph.nodes:
@@ -19,20 +19,11 @@ import onnx
import onnx_graphsurgeon as gs
import torch
from _test_utils.onnx.lib_test_models import SimpleMLP, export_as_onnx
from _test_utils.onnx.quantization.utils import assert_nodes_are_quantized
import modelopt.onnx.quantization as moq
def assert_nodes_are_quantized(nodes):
for node in nodes:
for inp_idx, inp in enumerate(node.inputs):
if isinstance(inp, gs.Variable):
assert node.i(inp_idx).op == "DequantizeLinear", (
f"Input '{inp.name}' of node '{node.name}' is not quantized but should be!"
)
return True
def test_fp8(tmp_path):
model_torch = SimpleMLP()
input_tensor = torch.randn(2, 16, 16)
@@ -29,34 +29,32 @@ import tempfile
import pytest
import torch
from _test_utils.torch.transformers_models import get_tiny_mixtral
from safetensors import safe_open
import modelopt.torch.quantization as mtq
from modelopt.torch.export import export_hf_checkpoint
pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="requires a GPU")
_SCALE_SUFFIXES = (".weight_scale", ".weight_scale_2", ".weight_scale_inv", ".input_scale")
# Wider than the shared defaults: NVFP4 needs weight dims that are a multiple of the block size.
_MIXTRAL_KWARGS = {
"hidden_size": 64,
"intermediate_size": 128,
"vocab_size": 320,
"max_position_embeddings": 128,
}
def _tiny_mixtral_config():
from transformers import MixtralConfig
cfg = MixtralConfig(
hidden_size=64,
intermediate_size=128,
num_hidden_layers=2,
num_attention_heads=4,
num_key_value_heads=2,
num_local_experts=4,
num_experts_per_tok=2,
vocab_size=320,
max_position_embeddings=128,
)
cfg.architectures = ["MixtralForCausalLM"]
return cfg
def _tiny_mixtral():
model = get_tiny_mixtral(**_MIXTRAL_KWARGS)
model.config.architectures = ["MixtralForCausalLM"]
return model
def test_export_tensor_names_match_hub_after_conversion_reverse():
pytest.importorskip("transformers")
from transformers import MixtralForCausalLM
try:
from transformers.conversion_mapping import get_checkpoint_conversion_mapping
from transformers.core_model_loading import revert_weight_conversion
@@ -65,19 +63,14 @@ def test_export_tensor_names_match_hub_after_conversion_reverse():
if not get_checkpoint_conversion_mapping("mixtral"):
pytest.skip("transformers build has no mixtral conversion_mapping")
import modelopt.torch.quantization as mtq
from modelopt.torch.export import export_hf_checkpoint
cfg = _tiny_mixtral_config()
# Canonical hub names: transformers' own reverse on the unquantized reference.
ref = MixtralForCausalLM(cfg)
ref = _tiny_mixtral()
hub_names = set(revert_weight_conversion(ref, ref.state_dict()).keys())
# sanity: reference really is fused/renamed in memory
assert any(".block_sparse_moe.experts.0.w1.weight" in n for n in hub_names)
model = MixtralForCausalLM(cfg).to("cuda", torch.bfloat16).eval()
ids = torch.randint(0, cfg.vocab_size, (2, 16), device="cuda")
model = _tiny_mixtral().to("cuda", torch.bfloat16).eval()
ids = torch.randint(0, model.config.vocab_size, (2, 16), device="cuda")
def forward_loop(m):
for _ in range(4):
@@ -90,8 +83,6 @@ def test_export_tensor_names_match_hub_after_conversion_reverse():
export_hf_checkpoint(model, export_dir=export_dir)
exported = set()
for f in glob.glob(os.path.join(export_dir, "*.safetensors")):
from safetensors import safe_open
with safe_open(f, framework="pt") as sf:
exported.update(sf.keys())
+19 -22
View File
@@ -15,24 +15,20 @@
import pytest
import torch
from transformers import LlamaConfig, LlamaForCausalLM, Qwen3MoeConfig, Qwen3MoeForCausalLM
from _test_utils.torch.transformers_models import get_tiny_llama, get_tiny_qwen3_moe
import modelopt.torch.quantization as mtq
from modelopt.torch.export.quant_utils import fuse_prequant_to_linear
def get_tiny_llama(attention_heads=4, key_value_heads=4):
"""Create a tiny Llama model for testing."""
config = LlamaConfig(
hidden_size=64,
intermediate_size=128,
num_hidden_layers=2,
num_attention_heads=attention_heads,
num_key_value_heads=key_value_heads,
max_position_embeddings=128,
vocab_size=256,
)
return LlamaForCausalLM(config)
# Wider than the shared defaults (AWQ block sizes need larger weight dims) and fp32 rather than
# the shared bf16 default, which the post-fusion allclose tolerances below are calibrated for.
_LLAMA_KWARGS = {
"dtype": torch.float32,
"hidden_size": 64,
"intermediate_size": 128,
"max_position_embeddings": 128,
"vocab_size": 256,
}
@pytest.mark.parametrize(
@@ -52,7 +48,12 @@ def get_tiny_llama(attention_heads=4, key_value_heads=4):
)
def test_pattern_fuse_prequant(quant_config, attention_kv_heads_pair):
"""Test pattern_fuse_prequant on modules from a tiny Llama model."""
model = get_tiny_llama(attention_kv_heads_pair[0], attention_kv_heads_pair[1]).to("cuda")
num_attention_heads, num_key_value_heads = attention_kv_heads_pair
model = get_tiny_llama(
**_LLAMA_KWARGS,
num_attention_heads=num_attention_heads,
num_key_value_heads=num_key_value_heads,
).to("cuda")
# Quantize the model
dummy_input = torch.randint(0, 256, (1, 16), device="cuda")
@@ -106,21 +107,17 @@ def test_pattern_fuse_prequant(quant_config, attention_kv_heads_pair):
def test_pattern_fuse_prequant_moe(quant_config):
"""Test pattern_fuse_prequant on Qwen3 MoE sparse MLP."""
# Create a tiny Qwen3MoE model for testing
config = Qwen3MoeConfig(
model = get_tiny_qwen3_moe(
dtype=torch.float32,
hidden_size=128,
intermediate_size=256,
moe_intermediate_size=256,
num_hidden_layers=2,
num_attention_heads=4,
num_key_value_heads=4,
num_experts=4,
num_experts_per_tok=2,
max_position_embeddings=128,
vocab_size=256,
shared_expert_intermediate_size=256,
)
model = Qwen3MoeForCausalLM(config).to("cuda")
).to("cuda")
# Quantize the model
dummy_input = torch.randint(0, 256, (1, 16), device="cuda")
@@ -12,15 +12,14 @@
# 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 copy
from copy import deepcopy
import pytest
import torch
import transformers
from _test_utils.torch.quantization.offload import make_cpu_offloaded_model, make_layerwise_cfg
from _test_utils.torch.transformers_models import create_tiny_llama_dir, create_tiny_qwen3_moe_dir
from accelerate import init_empty_weights, load_checkpoint_and_dispatch
from transformers import AutoConfig, AutoModelForCausalLM
from transformers import AutoModelForCausalLM
import modelopt.torch.quantization as mtq
from modelopt.torch.export import export_hf_vllm_fq_checkpoint
@@ -130,36 +129,6 @@ def _test_hf_vllm_export(tmp_path, quant_cfg, model_dir):
assert any("_amax" in k for k in state), f"input quantizer {name} should preserve _amax"
def _make_cpu_offloaded_model(tmp_path, num_hidden_layers=3):
"""Create a tiny LLaMA model with layer 0 offloaded to CPU via accelerate."""
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_hidden_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
with init_empty_weights():
model = AutoModelForCausalLM.from_config(config)
device_map = {
n: 0
for n, m in model.named_modules()
if "layers" not in n or n.split("layers.")[-1].isdigit()
}
device_map["model.layers.0"] = "cpu"
model = load_checkpoint_and_dispatch(model, tiny_llama_dir, device_map=device_map)
return model, config, tiny_llama_dir
def _make_layerwise_cfg(base_cfg):
"""Add layerwise=True to a quant config's algorithm field."""
cfg = copy.deepcopy(base_cfg)
algo = cfg.get("algorithm", "max")
if isinstance(algo, str):
cfg["algorithm"] = {"method": algo, "layerwise": True}
else:
algo["layerwise"] = True
return cfg
@pytest.mark.parametrize("quant_cfg", [mtq.FP8_DEFAULT_CFG])
def test_hf_vllm_export_offload(tmp_path, quant_cfg):
"""Verifies the inplace_mem_efficient=True path mutates offloaded weights in place
@@ -169,12 +138,12 @@ def test_hf_vllm_export_offload(tmp_path, quant_cfg):
"""
num_hidden_layers = 3
model, _config, _tiny_llama_dir = _make_cpu_offloaded_model(
model, _config, _tiny_llama_dir, _inputs = make_cpu_offloaded_model(
tmp_path / "offloaded", num_hidden_layers=num_hidden_layers
)
model.eval()
seq_cfg = _make_layerwise_cfg(quant_cfg)
seq_cfg = make_layerwise_cfg(quant_cfg)
def forward_loop(model):
input_ids = torch.randint(0, model.config.vocab_size, (1, 128)).cuda()
+5 -3
View File
@@ -20,6 +20,8 @@ from pathlib import Path
import pytest
import torch
import torch.nn.functional as F
from _test_utils.fs_utils import assert_unmodified_tree
from _test_utils.torch.transformers_models import create_tiny_llama_dir
_KERNELS_DIR = Path(__file__).parent
@@ -80,9 +82,7 @@ def sdpa_reference(q, k, v, b_start_loc, b_seq_len, is_causal=True):
@pytest.fixture(scope="module")
def tiny_llama_dir(tmp_path_factory):
"""Tiny Llama: 2 layers, 64 hidden, 4 q-heads, 2 kv-heads, head_dim=16."""
from _test_utils.torch.transformers_models import create_tiny_llama_dir
return create_tiny_llama_dir(
model_dir = create_tiny_llama_dir(
tmp_path_factory.mktemp("tiny_llama"),
with_tokenizer=True,
num_hidden_layers=2,
@@ -92,3 +92,5 @@ def tiny_llama_dir(tmp_path_factory):
intermediate_size=64,
max_position_embeddings=64,
)
with assert_unmodified_tree(model_dir) as path:
yield path
@@ -18,8 +18,6 @@
import pytest
import torch
diffusers = pytest.importorskip("diffusers")
from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_KERNEL_AVAILABLE
from modelopt.torch.kernels.sparsity.attention import diffusers_triton_attention as diffusers_mod
from modelopt.torch.kernels.sparsity.attention import ltx_triton_attention as ltx_mod
@@ -17,10 +17,21 @@ import copy
import json
import os
import shutil
from contextlib import nullcontext
import pytest
import torch
import torch.nn as nn
from _test_utils.torch.quantization.offload import (
make_cpu_offloaded_model,
make_layerwise_cfg,
make_layerwise_checkpoint_cfg,
make_tiny_llama_and_inputs,
)
from _test_utils.torch.quantization.quant_utils import (
assert_nvfp4_static_amaxes_fp32,
nvfp4_static_amax_dtypes,
)
from _test_utils.torch.transformers_models import create_tiny_llama_dir
from accelerate import init_empty_weights, load_checkpoint_and_dispatch
from accelerate.hooks import AlignDevicesHook, add_hook_to_module
@@ -28,7 +39,6 @@ from transformers import AutoConfig, AutoModelForCausalLM
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.extensions import get_cuda_ext_mx
from modelopt.torch.quantization.nn import TensorQuantizer
from modelopt.torch.quantization.utils import (
enable_weight_access_and_writeback,
is_quantized_linear,
@@ -60,26 +70,6 @@ NVFP4_WEIGHT_MSE_FP8_SWEEP_CFG = {
}
def _nvfp4_static_amax_dtypes(model):
amax_dtypes = {}
for name, module in model.named_modules():
if (
isinstance(module, TensorQuantizer)
and module.is_nvfp4_static
and module.amax is not None
):
amax_dtypes[name] = module.amax.dtype
return amax_dtypes
def _assert_nvfp4_static_amaxes_fp32(amax_dtypes, model_dtype, label):
assert amax_dtypes, f"{label}: expected NVFP4 static amaxes for model dtype {model_dtype}"
assert all(amax_dtype == torch.float32 for amax_dtype in amax_dtypes.values()), (
f"{label}: expected all NVFP4 static amaxes to be fp32 for model dtype {model_dtype}, "
f"got {amax_dtypes}"
)
@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
def test_transformers_mse_calibrate_fp32_amax_save_restore(tmp_path, dtype):
if get_cuda_ext_mx() is None:
@@ -91,8 +81,8 @@ def test_transformers_mse_calibrate_fp32_amax_save_restore(tmp_path, dtype):
cfg = copy.deepcopy(NVFP4_WEIGHT_MSE_FP8_SWEEP_CFG)
mtq.quantize(model, cfg, lambda model: model(input_ids))
amax_dtypes = _nvfp4_static_amax_dtypes(model)
_assert_nvfp4_static_amaxes_fp32(amax_dtypes, dtype, "mse calibrated")
amax_dtypes = nvfp4_static_amax_dtypes(model)
assert_nvfp4_static_amaxes_fp32(amax_dtypes, dtype, "mse calibrated")
with torch.no_grad():
output = model(input_ids).logits.detach().clone()
@@ -102,8 +92,8 @@ def test_transformers_mse_calibrate_fp32_amax_save_restore(tmp_path, dtype):
model.save_pretrained(ckpt_path)
assert os.path.exists(ckpt_path / "modelopt_state.pth")
restored_model = AutoModelForCausalLM.from_pretrained(ckpt_path, torch_dtype=dtype).cuda()
restored_amax_dtypes = _nvfp4_static_amax_dtypes(restored_model)
_assert_nvfp4_static_amaxes_fp32(restored_amax_dtypes, dtype, "restored")
restored_amax_dtypes = nvfp4_static_amax_dtypes(restored_model)
assert_nvfp4_static_amaxes_fp32(restored_amax_dtypes, dtype, "restored")
with torch.no_grad():
restored_output = restored_model(input_ids).logits.detach().clone()
@@ -150,61 +140,21 @@ def test_cpu_offloaded_tinyllama(tmp_path):
assert torch.allclose(output_ref.logits, output_test.logits)
def _make_cpu_offloaded_model(tmp_path, num_hidden_layers=3):
"""Create a tiny LLaMA model with layer 0 offloaded to CPU via accelerate."""
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_hidden_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
with init_empty_weights():
model = AutoModelForCausalLM.from_config(config)
device_map = {
n: 0
for n, m in model.named_modules()
if "layers" not in n or n.split("layers.")[-1].isdigit()
}
device_map["model.layers.0"] = "cpu"
model = load_checkpoint_and_dispatch(model, tiny_llama_dir, device_map=device_map)
inputs = torch.randint(0, config.vocab_size, (1, 4)).cuda()
return model, config, tiny_llama_dir, inputs
def _make_layerwise_cfg(base_cfg):
"""Add layerwise=True to a quant config's algorithm field."""
cfg = copy.deepcopy(base_cfg)
algo = cfg.get("algorithm", "max")
if isinstance(algo, str):
cfg["algorithm"] = {"method": algo, "layerwise": True}
else:
algo["layerwise"] = True
return cfg
def _make_layerwise_checkpoint_cfg(base_cfg, checkpoint_dir):
"""Add layerwise=True and layerwise_checkpoint_dir to a quant config's algorithm field."""
cfg = _make_layerwise_cfg(base_cfg)
cfg["algorithm"]["layerwise_checkpoint_dir"] = checkpoint_dir
return cfg
@pytest.mark.parametrize("use_checkpoint", [False, True], ids=["no_ckpt", "ckpt"])
def test_layerwise_calibrate_cpu_offloaded(tmp_path, use_checkpoint):
"""Layerwise calibration on CPU-offloaded model matches GPU-only reference."""
quant_cfg = mtq.NVFP4_AWQ_LITE_CFG
num_layers = 3
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
inputs = torch.randint(0, config.vocab_size, (1, 4)).cuda()
tiny_llama_dir, config, inputs = make_tiny_llama_and_inputs(tmp_path, num_layers)
if use_checkpoint:
ckpt_dir = str(tmp_path / "seq_ckpt")
seq_cfg = _make_layerwise_checkpoint_cfg(quant_cfg, ckpt_dir)
seq_cfg = make_layerwise_checkpoint_cfg(quant_cfg, ckpt_dir)
else:
seq_cfg = _make_layerwise_cfg(quant_cfg)
seq_cfg = make_layerwise_cfg(quant_cfg)
# Reference: GPU-only model with layerwise calibration
ref_cfg = _make_layerwise_cfg(quant_cfg)
ref_cfg = make_layerwise_cfg(quant_cfg)
model_ref = AutoModelForCausalLM.from_pretrained(
tiny_llama_dir, torch_dtype=config.torch_dtype
).cuda()
@@ -247,12 +197,10 @@ def test_sequential_checkpoint_resume_cpu_offloaded(tmp_path):
"""Resume from a partial checkpoint on a CPU-offloaded model matches a full run."""
quant_cfg = mtq.NVFP4_AWQ_LITE_CFG
num_layers = 3
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
inputs = torch.randint(0, config.vocab_size, (1, 4)).cuda()
tiny_llama_dir, config, inputs = make_tiny_llama_and_inputs(tmp_path, num_layers)
ckpt_dir = str(tmp_path / "seq_ckpt")
seq_ckpt_cfg = _make_layerwise_checkpoint_cfg(quant_cfg, ckpt_dir)
seq_ckpt_cfg = make_layerwise_checkpoint_cfg(quant_cfg, ckpt_dir)
# Full reference run with checkpointing
with init_empty_weights():
@@ -300,12 +248,10 @@ def test_sequential_checkpoint_resume_cpu_offloaded(tmp_path):
def test_sequential_checkpoint_resume_multi_offload(tmp_path):
"""Resume with multiple layers offloaded exercises per-layer device resolution."""
num_layers = 3
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
inputs = torch.randint(0, config.vocab_size, (1, 4)).cuda()
tiny_llama_dir, config, inputs = make_tiny_llama_and_inputs(tmp_path, num_layers)
ckpt_dir = str(tmp_path / "seq_ckpt")
seq_ckpt_cfg = _make_layerwise_checkpoint_cfg(mtq.INT4_AWQ_CFG, ckpt_dir)
seq_ckpt_cfg = make_layerwise_checkpoint_cfg(mtq.INT4_AWQ_CFG, ckpt_dir)
def _make_multi_offload_model():
with init_empty_weights():
@@ -362,9 +308,7 @@ def _make_gptq_sequential_checkpoint_cfg(base_cfg, checkpoint_dir):
def test_sequential_gptq_cpu_offloaded(tmp_path, use_checkpoint):
"""Sequential GPTQ (weight-modifying) on CPU-offloaded model matches GPU-only reference."""
num_layers = 3
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
inputs = torch.randint(0, config.vocab_size, (1, 4)).cuda()
tiny_llama_dir, config, inputs = make_tiny_llama_and_inputs(tmp_path, num_layers)
if use_checkpoint:
ckpt_dir = str(tmp_path / "gptq_ckpt")
@@ -381,7 +325,7 @@ def test_sequential_gptq_cpu_offloaded(tmp_path, use_checkpoint):
output_ref = model_ref(inputs)
# Test: CPU-offloaded model
model, _, _, _ = _make_cpu_offloaded_model(tmp_path / "offloaded", num_hidden_layers=num_layers)
model, _, _, _ = make_cpu_offloaded_model(tmp_path / "offloaded", num_hidden_layers=num_layers)
mtq.quantize(model, seq_cfg, lambda model: model(inputs))
output_test = model(inputs)
@@ -398,9 +342,7 @@ def test_sequential_gptq_cpu_offloaded(tmp_path, use_checkpoint):
def test_sequential_gptq_checkpoint_resume_cpu_offloaded(tmp_path):
"""GPTQ checkpoint resume with CPU offloading restores modified weights correctly."""
num_layers = 3
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
inputs = torch.randint(0, config.vocab_size, (1, 4)).cuda()
tiny_llama_dir, config, inputs = make_tiny_llama_and_inputs(tmp_path, num_layers)
ckpt_dir = str(tmp_path / "gptq_ckpt")
seq_ckpt_cfg = _make_gptq_sequential_checkpoint_cfg(mtq.NVFP4_AWQ_LITE_CFG, ckpt_dir)
@@ -474,8 +416,6 @@ class _TupleUnpackingModel(torch.nn.Module):
def test_skip_dummy_has_no_hf_hook(monkeypatch):
"""Dummies must not carry _hf_hook from the original layer."""
from contextlib import nullcontext
monkeypatch.setattr(
LayerActivationCollector,
"_decoder_layer_support",
@@ -512,8 +452,6 @@ def test_skip_dummy_has_no_hf_hook(monkeypatch):
def _assert_persistent_materialization_bypasses_top_hook(layer):
from modelopt.torch.quantization.utils import persistent_materialization
assert hasattr(layer, "_hf_hook")
original_old_forward = layer._old_forward
@@ -533,7 +471,7 @@ def _assert_persistent_materialization_bypasses_top_hook(layer):
def test_persistent_materialization_cpu_offloaded(tmp_path):
"""persistent_materialization keeps CPU-offloaded weights on GPU and writes back modifications."""
model, config, _, inputs = _make_cpu_offloaded_model(tmp_path)
model, config, _, inputs = make_cpu_offloaded_model(tmp_path)
offloaded_layer = model.model.layers[0]
# Verify offloaded (meta device)
@@ -695,18 +633,16 @@ def test_layerwise_calibrate_disk_offloaded(tmp_path, use_checkpoint):
"""Layerwise calibration on disk-offloaded model matches GPU-only reference."""
quant_cfg = mtq.NVFP4_AWQ_LITE_CFG
num_layers = 3
tiny_llama_dir = create_tiny_llama_dir(tmp_path, num_hidden_layers=num_layers)
config = AutoConfig.from_pretrained(tiny_llama_dir)
inputs = torch.randint(0, config.vocab_size, (1, 4)).cuda()
tiny_llama_dir, config, inputs = make_tiny_llama_and_inputs(tmp_path, num_layers)
if use_checkpoint:
ckpt_dir = str(tmp_path / "seq_ckpt")
seq_cfg = _make_layerwise_checkpoint_cfg(quant_cfg, ckpt_dir)
seq_cfg = make_layerwise_checkpoint_cfg(quant_cfg, ckpt_dir)
else:
seq_cfg = _make_layerwise_cfg(quant_cfg)
seq_cfg = make_layerwise_cfg(quant_cfg)
# Reference: GPU-only model with layerwise calibration
ref_cfg = _make_layerwise_cfg(quant_cfg)
ref_cfg = make_layerwise_cfg(quant_cfg)
model_ref = AutoModelForCausalLM.from_pretrained(
tiny_llama_dir, torch_dtype=config.torch_dtype
).cuda()
@@ -23,6 +23,7 @@ import inspect
import pytest
import torch
from _test_utils.torch.quantization.attention import make_quant_attention
from transformers import LlamaConfig
from transformers.models.llama.modeling_llama import LlamaAttention
@@ -34,26 +35,13 @@ except ImportError:
from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_FA_AVAILABLE
from modelopt.torch.quantization.plugins.huggingface import _QuantAttention
pytest.importorskip("transformers")
def _make_quant_attention(hidden_size=128, num_q_heads=4, num_kv_heads=2):
config = LlamaConfig(
hidden_size=hidden_size,
num_attention_heads=num_q_heads,
num_key_value_heads=num_kv_heads,
)
quant_attention = _QuantAttention.convert(LlamaAttention(config, layer_idx=0))
quant_attention.config._attn_implementation = "sdpa"
return quant_attention
@pytest.mark.skipif(not TRITON_FA_AVAILABLE, reason="Triton attention kernel unavailable")
def test_p_qdq_fa():
"""FP8/NVFP4 p_bmm_quantizer runs on the built-in Triton kernel (no kitchen)."""
batch_size, num_q_heads, num_kv_heads, seqlen, head_dim = 2, 4, 2, 32, 64
quant_attention = _make_quant_attention(num_q_heads=num_q_heads, num_kv_heads=num_kv_heads)
quant_attention = make_quant_attention(num_q_heads=num_q_heads, num_kv_heads=num_kv_heads)
for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"):
getattr(quant_attention, name).disable()
@@ -113,7 +101,7 @@ def test_p_qdq_unsupported_cases_raise():
"""The Triton qdq dispatch rejects attention semantics the kernel cannot honor."""
batch_size, num_q_heads, num_kv_heads, seqlen, head_dim = 2, 4, 2, 32, 64
quant_attention = _make_quant_attention(num_q_heads=num_q_heads, num_kv_heads=num_kv_heads)
quant_attention = make_quant_attention(num_q_heads=num_q_heads, num_kv_heads=num_kv_heads)
for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"):
getattr(quant_attention, name).disable()
quant_attention.p_bmm_quantizer.num_bits = (4, 3) # FP8 mode
@@ -172,7 +160,7 @@ def test_p_qdq_non_causal_falls_back_to_eager():
instead of raising -- keeping the softmax-P quant in an export-traceable graph."""
batch_size, num_q_heads, num_kv_heads, seqlen, head_dim = 2, 4, 2, 32, 64
quant_attention = _make_quant_attention(num_q_heads=num_q_heads, num_kv_heads=num_kv_heads)
quant_attention = make_quant_attention(num_q_heads=num_q_heads, num_kv_heads=num_kv_heads)
for name in ("q_bmm_quantizer", "k_bmm_quantizer", "v_bmm_quantizer"):
getattr(quant_attention, name).disable()
+3 -10
View File
@@ -24,9 +24,11 @@ import torch
import torch.nn as nn
from _test_utils.torch.distributed.utils import synchronize_state_dict
from torch.distributed._composable.fsdp.fully_shard import fully_shard
from torch.distributed.fsdp import CPUOffloadPolicy
from torch.distributed.tensor import DTensor
import modelopt.torch.quantization as mtq
import modelopt.torch.quantization.model_calib as model_calib
from modelopt.torch.opt.dynamic import _pytorch_managed
from modelopt.torch.quantization.nn import StaticBlockScaleQuantizer, TensorQuantizer
from modelopt.torch.quantization.utils import (
@@ -34,6 +36,7 @@ from modelopt.torch.quantization.utils import (
persistent_materialization,
)
from modelopt.torch.quantization.utils.layerwise_calib import LayerActivationCollector
from modelopt.torch.utils.dataset_utils import _forward_loop
def _test_fsdp2_simple_linear(rank, size):
@@ -214,8 +217,6 @@ class _SimpleTransformerModel(nn.Module):
def _test_layerwise_calibrate_fsdp2(rank, size):
"""Layerwise calibration on FSDP2-wrapped model matches non-FSDP reference."""
import modelopt.torch.quantization.model_calib as model_calib
dim = 32
torch.manual_seed(1)
model = _SimpleTransformerModel(n_layers=3, dim=dim).cuda()
@@ -342,8 +343,6 @@ def _test_writeback_root_unwrapped(rank, size):
(``fsdp2_wrap`` now defaults to ``shard_root=True``, wrapping the root too). Regression guard
for the old ``isinstance(root_model, FSDPModule)`` assert that wrongly required a wrapped root.
"""
from modelopt.torch.quantization.utils import enable_weight_access_and_writeback
dim = 32
torch.manual_seed(1)
# Root is a plain container; model[0] stands in for a decoder layer.
@@ -385,10 +384,6 @@ def _test_writeback_cpu_offload(rank, size):
so the helper mirrors it to GPU for in-context mutation and must copy
modifications back to the CPU shard on exit.
"""
from torch.distributed.fsdp import CPUOffloadPolicy
from modelopt.torch.quantization.utils import enable_weight_access_and_writeback
dim = 32
torch.manual_seed(1)
model = nn.Sequential(nn.Sequential(nn.Linear(dim, dim), nn.Linear(dim, dim))).cuda(rank)
@@ -440,8 +435,6 @@ def _test_sharded_root_calibration(rank, size):
unshard embed/norm for the forward and reshard them after — no manual materialization.
With the old ``model.forward`` bypass this hit ``aten.embedding: mixed Tensor and DTensor``.
"""
from modelopt.torch.utils.dataset_utils import _forward_loop
dim = 32
torch.manual_seed(1)
model = _EmbedRootModel(dim=dim).cuda(rank)
@@ -20,6 +20,11 @@ import torch
from torch import nn
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.nn.modules.tensor_quantizer import (
StaticBlockScaleQuantizer,
TensorQuantizer,
)
from modelopt.torch.quantization.tensor_quant import fp4_cast_ste
NVFP4_LSQ_POST_MSE_CFG = {
"quant_cfg": {
@@ -149,11 +154,6 @@ def test_lsq_quantize_e2e(config):
def test_lsq_fp4_fake_quantize_differentiable():
"""Test that _fake_quantize in FP4 LSQ mode is differentiable."""
from modelopt.torch.quantization.nn.modules.tensor_quantizer import (
StaticBlockScaleQuantizer,
TensorQuantizer,
)
device = torch.device("cuda")
tq = TensorQuantizer()
tq._num_bits = (2, 1)
@@ -184,8 +184,6 @@ def test_lsq_fp4_fake_quantize_differentiable():
def test_lsq_fp4_cast_ste():
"""Test fp4_cast_ste on GPU."""
from modelopt.torch.quantization.tensor_quant import fp4_cast_ste
device = torch.device("cuda")
x = torch.tensor([[-3.0, 1.5, 0.0, 6.0, -6.0, 0.5, -0.5, 2.0]], device=device)
x.requires_grad_(True)
@@ -29,6 +29,10 @@ from contextlib import contextmanager
import pytest
import torch
from _test_utils.torch.quantization.models import SimpleLinear
from _test_utils.torch.quantization.quant_utils import (
assert_nvfp4_static_amaxes_fp32,
nvfp4_static_amax_dtypes,
)
from conftest import requires_triton
import modelopt.torch.opt as mto
@@ -40,7 +44,6 @@ from modelopt.torch.kernels.quantization.gemm import (
from modelopt.torch.quantization.calib import NVFP4MSECalibrator
from modelopt.torch.quantization.extensions import get_cuda_ext_mx
from modelopt.torch.quantization.model_calib import _LocalHessianAccumulator
from modelopt.torch.quantization.nn import TensorQuantizer
from modelopt.torch.quantization.tensor_quant import static_blockwise_fp4_fake_quant
from modelopt.torch.quantization.utils.numeric_utils import E4M3_MAX
@@ -81,26 +84,6 @@ def _make_calibrator(per_block_amax, global_amax):
)
def _nvfp4_static_amax_dtypes(model):
amax_dtypes = {}
for name, module in model.named_modules():
if (
isinstance(module, TensorQuantizer)
and module.is_nvfp4_static
and module.amax is not None
):
amax_dtypes[name] = module.amax.dtype
return amax_dtypes
def _assert_nvfp4_static_amaxes_fp32(amax_dtypes, model_dtype, label):
assert amax_dtypes, f"{label}: expected NVFP4 static amaxes for model dtype {model_dtype}"
assert all(amax_dtype == torch.float32 for amax_dtype in amax_dtypes.values()), (
f"{label}: expected all NVFP4 static amaxes to be fp32 for model dtype {model_dtype}, "
f"got {amax_dtypes}"
)
def _run_reference(x, per_block_amax, global_amax):
with _force_sweep_path(triton_enabled=False):
cal = _make_calibrator(per_block_amax, global_amax)
@@ -338,14 +321,14 @@ def test_mse_calibrate_end_to_end(monkeypatch, tmp_path, dtype):
m(batch)
mtq.quantize(model, cfg, forward_loop=forward_loop)
amax_dtypes = _nvfp4_static_amax_dtypes(model)
_assert_nvfp4_static_amaxes_fp32(amax_dtypes, dtype, label)
amax_dtypes = nvfp4_static_amax_dtypes(model)
assert_nvfp4_static_amaxes_fp32(amax_dtypes, dtype, label)
ckpt_path = tmp_path / f"mse_calibrate_{label}_{str(dtype).rpartition('.')[-1]}.pt"
mto.save(model, ckpt_path)
restored_model = mto.restore(SimpleLinear(dtype=dtype).cuda(), ckpt_path)
restored_amax_dtypes = _nvfp4_static_amax_dtypes(restored_model)
_assert_nvfp4_static_amaxes_fp32(restored_amax_dtypes, dtype, f"{label} restored")
restored_amax_dtypes = nvfp4_static_amax_dtypes(restored_model)
assert_nvfp4_static_amaxes_fp32(restored_amax_dtypes, dtype, f"{label} restored")
# Run a deterministic input through and snapshot the output.
torch.manual_seed(1)
@@ -367,10 +350,10 @@ def test_mse_calibrate_end_to_end(monkeypatch, tmp_path, dtype):
# under the same seed) before we compare post-calibration outputs.
for name in w0:
assert torch.equal(w0[name], w1[name]), name
_assert_nvfp4_static_amaxes_fp32(dtypes_default, dtype, "fast")
_assert_nvfp4_static_amaxes_fp32(dtypes_optout, dtype, "reference")
_assert_nvfp4_static_amaxes_fp32(restored_dtypes_default, dtype, "fast restored")
_assert_nvfp4_static_amaxes_fp32(restored_dtypes_optout, dtype, "reference restored")
assert_nvfp4_static_amaxes_fp32(dtypes_default, dtype, "fast")
assert_nvfp4_static_amaxes_fp32(dtypes_optout, dtype, "reference")
assert_nvfp4_static_amaxes_fp32(restored_dtypes_default, dtype, "fast restored")
assert_nvfp4_static_amaxes_fp32(restored_dtypes_optout, dtype, "reference restored")
assert y_default.dtype == dtype
assert y_optout.dtype == dtype
assert torch.equal(y_default, y_optout)
@@ -31,9 +31,8 @@ pytestmark = [
pytest.mark.filterwarnings("ignore::DeprecationWarning"),
]
diffusers = pytest.importorskip("diffusers")
import numpy as np
from _test_utils.torch.diffusers_models import get_tiny_wan22_transformer
from diffusers import WanPipeline
import modelopt.torch.opt as mto
@@ -41,6 +40,9 @@ from modelopt.torch.kernels.common.attention import IS_AVAILABLE as TRITON_KERNE
if TRITON_KERNEL_AVAILABLE:
import modelopt.torch.sparsity.attention_sparsity as mtsa
from modelopt.torch.sparsity.attention_sparsity.methods.triton_skip_softmax import (
TritonSkipSoftmaxMethod,
)
from modelopt.torch.sparsity.attention_sparsity.sparse_attention import SparseAttentionModule
@@ -49,14 +51,6 @@ if TRITON_KERNEL_AVAILABLE:
# ---------------------------------------------------------------------------
@pytest.fixture(scope="module")
def tiny_wan22_path(tmp_path_factory):
"""Create and save a tiny Wan 2.2 pipeline to disk once per module."""
from _test_utils.torch.diffusers_models import create_tiny_wan22_pipeline_dir
return str(create_tiny_wan22_pipeline_dir(tmp_path_factory.mktemp("tiny_wan22")))
@pytest.fixture
def tiny_wan22_pipe(tiny_wan22_path):
"""Load a fresh copy of the tiny Wan 2.2 pipeline on CUDA (per test)."""
@@ -168,10 +162,6 @@ class TestWan22PipelineE2E:
def test_measure_sparsity_counts_accumulate(self, tiny_wan22_pipe):
"""measure_sparsity=True + a permissive threshold → nonzero sparsity counters."""
from modelopt.torch.sparsity.attention_sparsity.methods.triton_skip_softmax import (
TritonSkipSoftmaxMethod,
)
_sparsify_both_transformers(tiny_wan22_pipe, _skip_softmax_cfg(threshold=0.25))
# Enable measurement + reset counters on every sparse module
@@ -204,8 +194,6 @@ class TestWan22PipelineE2E:
``attn2`` modules keep the default method. The restored model must show the
identical (module_name → method) mapping.
"""
from _test_utils.torch.diffusers_models import get_tiny_wan22_transformer
_sparsify_both_transformers(tiny_wan22_pipe, _skip_softmax_cfg())
state = mto.modelopt_state(tiny_wan22_pipe.transformer)
@@ -247,10 +235,6 @@ class TestWan22Calibration:
def test_calibration_collects_stats_per_module(self, tiny_wan22_pipe):
"""A forward pass under calibration_mode populates per-module _last_stats."""
from modelopt.torch.sparsity.attention_sparsity.methods.triton_skip_softmax import (
TritonSkipSoftmaxMethod,
)
_sparsify_both_transformers(tiny_wan22_pipe, _skip_softmax_cfg())
threshold_trials = [1e-3, 1e-2, 1e-1]
@@ -18,37 +18,23 @@
These tests require a CUDA GPU. CPU-only tests are in tests/unit/.
"""
from copy import deepcopy
import pytest
import torch
from _test_utils.torch.speculative.dflash import get_dflash_config
from _test_utils.torch.transformers_models import get_tiny_llama
import modelopt.torch.speculative as mtsp
from modelopt.torch.speculative.config import DFLASH_DEFAULT_CFG
BLOCK_SIZE = 4
NUM_DRAFT_LAYERS = 2
SEQ_LEN = 16 # must be multiple of BLOCK_SIZE
def _get_dflash_config(block_size=BLOCK_SIZE, num_layers=NUM_DRAFT_LAYERS):
"""Create a DFlash config for testing."""
config = deepcopy(DFLASH_DEFAULT_CFG["config"])
config["dflash_block_size"] = block_size
config["dflash_use_torch_compile"] = False
config["dflash_mask_token_id"] = 0
config["dflash_architecture_config"] = {
"num_hidden_layers": num_layers,
}
return config
@pytest.fixture
def dflash_model():
"""Create a tiny DFlash model on GPU."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
model = model.cuda()
return model
@@ -120,7 +106,7 @@ class TestDFlashTrainingForwardGPU:
def model(self):
"""Create a tiny DFlash model in training mode on GPU."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
model = model.cuda()
model.train()
@@ -206,7 +192,7 @@ class TestDFlashOfflineForwardGPU:
"""Create a tiny DFlash model with dflash_offline=True on GPU."""
model = get_tiny_llama(num_hidden_layers=self.NUM_BASE_LAYERS)
model.config.num_orig_hidden_layers = self.NUM_BASE_LAYERS
config = _get_dflash_config()
config = get_dflash_config()
config["dflash_offline"] = True
mtsp.convert(model, [("dflash", config)])
model = model.cuda()
+23 -51
View File
@@ -17,21 +17,23 @@
import json
import os
import tempfile
from functools import partial
import pytest
import torch
import torch.distributed as dist
from _test_utils.torch.transformers_models import create_tiny_llama_dir
from torch.distributed.tensor import DTensor
from modelopt.torch.export.unified_export_hf import export_hf_checkpoint
from modelopt.torch.utils.distributed import broadcast_state_dict
from modelopt.torch.utils.plugins.model_load_utils import parallel_load_and_prepare_fsdp2
VOCAB_SIZE = 64
def _test_broadcast_state_dict_roundtrip(rank, size):
"""Round-trip from every rank as source (matches the per-layer rotation in the loader)."""
from modelopt.torch.utils.distributed import broadcast_state_dict
"""Round-trip from every rank as source, with a distinct payload per source rank."""
device = torch.device(f"cuda:{rank}")
# Distinct payload per source rank so a wrong-src result would fail content checks.
for source in range(size):
src_dict = {
"w": torch.full((2, 4), float(source)),
@@ -48,41 +50,8 @@ def test_broadcast_state_dict_roundtrip(dist_workers):
dist_workers.run(_test_broadcast_state_dict_roundtrip)
def _build_tiny_llama_checkpoint(path: str) -> None:
"""Write a tiny LlamaForCausalLM checkpoint (config + safetensors) to ``path``."""
from transformers import LlamaConfig, LlamaForCausalLM
config = LlamaConfig(
vocab_size=64,
hidden_size=32,
intermediate_size=64,
num_hidden_layers=2,
num_attention_heads=4,
max_position_embeddings=32,
torch_dtype="bfloat16",
)
model = LlamaForCausalLM(config).to(torch.bfloat16)
model.save_pretrained(path)
def _test_parallel_load_and_export(rank, size, cpu_offload):
"""Load a tiny Llama via the FSDP2 loader, forward, then export — config.architectures preserved.
Parametrized over ``cpu_offload`` to cover both shard placements:
- off: decoder DTensor shards on GPU, plain root on GPU.
- on: decoder DTensor shards on CPU (streamed per layer), root promoted to GPU
via ``_promote_non_dtensor_to_gpu``.
"""
from modelopt.torch.export.unified_export_hf import export_hf_checkpoint
from modelopt.torch.utils.plugins.model_load_utils import parallel_load_and_prepare_fsdp2
suffix = "offload" if cpu_offload else "noffload"
ckpt_dir = os.path.join(tempfile.gettempdir(), f"_test_parallel_load_{suffix}_{os.getpid()}")
if rank == 0:
os.makedirs(ckpt_dir, exist_ok=True)
_build_tiny_llama_checkpoint(ckpt_dir)
dist.barrier()
def _test_parallel_load_and_export(rank, size, ckpt_dir, export_dir, cpu_offload):
"""Load a tiny Llama via the FSDP2 loader, forward, then export."""
device = torch.device(f"cuda:{rank}")
model = parallel_load_and_prepare_fsdp2(
ckpt_dir,
@@ -105,17 +74,11 @@ def _test_parallel_load_and_export(rank, size, cpu_offload):
assert all(p.to_local().device.type == "cpu" for p in decoder_dtensors)
# Forward exercises FSDP2 hooks + (under cpu_offload) the per-layer CPU↔GPU stream.
input_ids = torch.randint(0, 64, (1, 8), device=device)
input_ids = torch.randint(0, VOCAB_SIZE, (1, 8), device=device)
out = model(input_ids=input_ids).logits
assert out.shape == (1, 8, 64)
assert out.shape == (1, 8, VOCAB_SIZE)
# Export and verify the saved config.json retains the original architectures.
export_dir = os.path.join(
tempfile.gettempdir(), f"_test_parallel_export_{suffix}_{os.getpid()}"
)
if rank == 0:
os.makedirs(export_dir, exist_ok=True)
dist.barrier()
export_hf_checkpoint(model, export_dir=export_dir, dtype=torch.bfloat16)
if rank == 0:
@@ -125,5 +88,14 @@ def _test_parallel_load_and_export(rank, size, cpu_offload):
@pytest.mark.parametrize("cpu_offload", [False, True])
def test_parallel_load_and_export(dist_workers, cpu_offload):
dist_workers.run(partial(_test_parallel_load_and_export, cpu_offload=cpu_offload))
def test_parallel_load_and_export(dist_workers, tmp_path, cpu_offload):
# Build the checkpoint once here (not inside the workers): every rank must see the same path.
ckpt_dir = create_tiny_llama_dir(tmp_path, vocab_size=VOCAB_SIZE)
dist_workers.run(
partial(
_test_parallel_load_and_export,
ckpt_dir=str(ckpt_dir),
export_dir=str(tmp_path / "export"),
cpu_offload=cpu_offload,
)
)
+4 -3
View File
@@ -16,10 +16,12 @@ import contextlib
import pytest
import torch
from _test_utils.fs_utils import assert_unmodified_tree
from _test_utils.torch.distributed.utils import DistributedWorkerPool
from _test_utils.torch.transformers_models import get_tiny_tokenizer
from megatron.core.parallel_state import destroy_model_parallel
import modelopt.torch.quantization.extensions as ext
import modelopt.torch.utils.distributed as dist
@@ -27,7 +29,8 @@ import modelopt.torch.utils.distributed as dist
def tiny_tokenizer_path(tmp_path_factory):
tokenizer_path = tmp_path_factory.mktemp("tiny_tokenizer")
get_tiny_tokenizer().save_pretrained(tokenizer_path)
return str(tokenizer_path)
with assert_unmodified_tree(tokenizer_path) as path:
yield str(path)
apex_destroy = None
@@ -48,8 +51,6 @@ def _prebuild_quant_cuda_extensions():
is not itself capped by a per-test timeout. Worker subprocesses then load the cached
.so from the shared ``TORCH_EXTENSIONS_DIR``.
"""
import modelopt.torch.quantization.extensions as ext
ext.precompile()
@@ -32,6 +32,7 @@ from safetensors import safe_open
from safetensors.torch import save_file
from transformers.models.qwen3_vl.modeling_qwen3_vl import Qwen3VLForConditionalGeneration
import modelopt.torch.export.unified_export_megatron as uem
import modelopt.torch.quantization as mtq
import modelopt.torch.speculative as mtsp
from modelopt.torch.export import KV_CACHE_FP8, export_mcore_gpt_to_hf, import_mcore_gpt_from_hf
@@ -623,8 +624,6 @@ def test_is_sidecar_writer_rank_pins_to_dp0_ep0(monkeypatch):
"""DP>1 fix predicate: only the DP0/EP0 rank among is_last_stage_main_rank writes
sidecar files. Guards the predicate used at three sites in save_pretrained.
"""
import modelopt.torch.export.unified_export_megatron as uem
# is_last_stage_main_rank=False is never a writer, regardless of DP/EP.
monkeypatch.setattr(uem, "get_data_parallel_rank", lambda: 0)
monkeypatch.setattr(uem, "get_expert_model_parallel_rank", lambda: 0)
@@ -561,6 +561,7 @@ def test_homogeneous_sharded_state_dict_hybrid(dist_workers, tmp_path, config):
mixed_block_size_config,
],
)
@skip_flaky_on_blackwell
def test_heterogenous_sharded_state_dict(dist_workers, tmp_path, config):
dist_workers.run(
partial(_test_sharded_state_dict, tmp_path, config, 256, None, False, False, {}),
@@ -18,6 +18,7 @@ import copy
import pytest
import torch
import torch.nn as nn
import transformer_engine as te
from _test_utils.torch.misc import set_seed
from _test_utils.torch.quantization.quantize_common import quantize_model_and_forward
@@ -26,8 +27,6 @@ import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.extensions import get_cuda_ext_mx
from modelopt.torch.quantization.nn import QuantModule
te = pytest.importorskip("transformer_engine")
class TELinear(nn.Module):
def __init__(self):
@@ -36,8 +36,6 @@ def test_export_compressed_nvfp4_weight_trtllm_scale():
``weight_quantizer._scale`` instead of the modelopt 2-D E4M3 layout. The export must
un-swizzle it; using it as-is would write a scale of raw byte values.
"""
pytest.importorskip("tensorrt_llm")
in_features = 256
calib = lambda x: x(torch.randn(1, 4, in_features).cuda().half()) # noqa: E731
@@ -15,6 +15,9 @@
import pytest
import torch
from _test_utils.torch.quantization.models import OneLayerLinear
from tensorrt_llm._torch.auto_deploy.utils.quantization_utils import (
cutlass_fp4_scale_to_modelopt_fp4_scale,
)
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.backends.utils import fp4_compatible
@@ -24,10 +27,6 @@ from modelopt.torch.quantization.qtensor import NVFP4QTensor
@pytest.mark.skipif(not fp4_compatible(), reason="FP4 is not supported on this GPU")
@pytest.mark.parametrize("shape", [(128, 64), (3, 16)])
def test_nvfp4_quantization(shape):
from tensorrt_llm._torch.auto_deploy.utils.quantization_utils import (
cutlass_fp4_scale_to_modelopt_fp4_scale,
)
block_sizes = {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}
weight = torch.randn(shape).to(torch.float16).cuda()
@@ -32,21 +32,12 @@ from _test_utils.onnx.lib_test_models import (
build_small_grouped_conv_model,
export_as_onnx,
)
from _test_utils.onnx.quantization.utils import assert_nodes_are_quantized
from modelopt.onnx.quantization.quantize import quantize
from modelopt.onnx.utils import get_opset_version, save_onnx
def assert_nodes_are_quantized(nodes):
for node in nodes:
for inp_idx, inp in enumerate(node.inputs):
if isinstance(inp, gs.Variable) and node.i(inp_idx).op != "Identity":
assert node.i(inp_idx).op == "DequantizeLinear", (
f"Input '{inp.name}' of node '{node.name}' is not quantized but should be!"
)
return True
def assert_nodes_are_not_quantized(nodes):
for node in nodes:
for inp_idx, inp in enumerate(node.inputs):
@@ -77,7 +68,7 @@ def test_bias_add_rule(tmp_path):
# Check that all Conv nodes are quantized
conv_nodes = [n for n in graph.nodes if n.op == "Conv"]
assert assert_nodes_are_quantized(conv_nodes)
assert assert_nodes_are_quantized(conv_nodes, ignore_identity_inputs=True)
# Check that all other nodes are not quantized
other_nodes = [
@@ -99,7 +90,7 @@ def _check_resnet_residual_connection(onnx_path):
# Check that all Conv nodes are quantized
conv_nodes = [n for n in graph.nodes if n.op == "Conv"]
assert assert_nodes_are_quantized(conv_nodes)
assert assert_nodes_are_quantized(conv_nodes, ignore_identity_inputs=True)
# Check that the left-side branch of Add contains a QDQ node
# In this case, this means that the inputs of Add should be DequantizeLinear and Conv.
@@ -148,7 +139,7 @@ def test_convtranspose_conv_residual_int8(tmp_path):
# Check that Conv and ConvTransposed are quantized
conv_nodes = [n for n in graph.nodes if "Conv" in n.op]
assert assert_nodes_are_quantized(conv_nodes)
assert assert_nodes_are_quantized(conv_nodes, ignore_identity_inputs=True)
# Check that only 1 input of Add is quantized
add_nodes = [n for n in graph.nodes if n.op == "Add"]
@@ -177,7 +168,7 @@ def test_conv_batchnorm_sig_mul_int8(tmp_path):
# Check that Conv and ConvTransposed are quantized
conv_nodes = [n for n in graph.nodes if "Conv" in n.op]
assert assert_nodes_are_quantized(conv_nodes)
assert assert_nodes_are_quantized(conv_nodes, ignore_identity_inputs=True)
# Check that only 1 input of Add is quantized
add_nodes = [n for n in graph.nodes if n.op == "Add"]
@@ -207,7 +198,7 @@ def test_conv_act_pool_int8(tmp_path, include_reshape_node):
# Check that Conv is quantized
conv_nodes = [n for n in graph.nodes if n.op == "Conv"]
assert assert_nodes_are_quantized(conv_nodes)
assert assert_nodes_are_quantized(conv_nodes, ignore_identity_inputs=True)
# Check that MaxPool is not quantized
pool_nodes = [n for n in graph.nodes if n.op == "MaxPool"]
@@ -233,7 +224,7 @@ def test_conv_isinf_int8(tmp_path):
# Check that Conv is quantized
conv_nodes = [n for n in graph.nodes if "Conv" in n.op]
assert assert_nodes_are_quantized(conv_nodes)
assert assert_nodes_are_quantized(conv_nodes, ignore_identity_inputs=True)
# Check that IsInf is running in the lowest supported precision:
# - FP32 if opset < 20, or
@@ -269,7 +260,7 @@ def test_conv_layernorm_quantization(tmp_path):
# Check that Conv nodes are quantized (inputs have Q/DQ)
conv_nodes = [n for n in graph.nodes if n.op == "Conv"]
assert assert_nodes_are_quantized(conv_nodes)
assert assert_nodes_are_quantized(conv_nodes, ignore_identity_inputs=True)
# Check that LayerNormalization has Q/DQ on its activation input
ln_nodes = [n for n in graph.nodes if n.op == "LayerNormalization"]
@@ -308,8 +299,8 @@ def test_target_dla_conv(tmp_path, target_dla):
mul_nodes = [n for n in graph.nodes if "Mul" in n.op]
if target_dla:
# Check that all Convs and Mul nodes are quantized
assert assert_nodes_are_quantized(conv_nodes)
assert assert_nodes_are_quantized(mul_nodes)
assert assert_nodes_are_quantized(conv_nodes, ignore_identity_inputs=True)
assert assert_nodes_are_quantized(mul_nodes, ignore_identity_inputs=True)
else:
# Check that only the 1st Conv is quantized
assert assert_nodes_are_quantized([conv_nodes[0]])
@@ -331,7 +322,7 @@ def test_target_dla_matmul(tmp_path, target_dla):
matmul_nodes = [n for n in graph.nodes if n.op == "MatMul"]
if target_dla:
# Check that MatMul is quantized
assert assert_nodes_are_quantized(matmul_nodes)
assert assert_nodes_are_quantized(matmul_nodes, ignore_identity_inputs=True)
else:
# GEMV detection excludes the MatMul (m=1) from quantization.
assert assert_nodes_are_not_quantized(matmul_nodes)
@@ -20,21 +20,12 @@ import onnx_graphsurgeon as gs
import pytest
import torch
from _test_utils.onnx.lib_test_models import SimpleMLP, export_as_onnx
from _test_utils.onnx.quantization.utils import assert_nodes_are_quantized
from onnxruntime.quantization.calibrate import CalibrationDataReader
import modelopt.onnx.quantization as moq
def assert_nodes_are_quantized(nodes):
for node in nodes:
for inp_idx, inp in enumerate(node.inputs):
if isinstance(inp, gs.Variable):
assert node.i(inp_idx).op == "DequantizeLinear", (
f"Input '{inp.name}' of node '{node.name}' is not quantized but should be!"
)
return True
def int8_test_helper(tmp_path, high_precision_dtype, **kwargs):
model_torch = SimpleMLP()
input_tensor = torch.randn(2, 16, 16)
+1 -10
View File
@@ -20,21 +20,12 @@ import onnx_graphsurgeon as gs
import pytest
import torch
from _test_utils.onnx.lib_test_models import SimpleMLP, export_as_onnx
from _test_utils.onnx.quantization.utils import assert_nodes_are_quantized
from modelopt.onnx.autocast import convert_to_mixed_precision
from modelopt.onnx.quantization import quantize
def assert_nodes_are_quantized(nodes):
for node in nodes:
for inp_idx, inp in enumerate(node.inputs):
if isinstance(inp, gs.Variable):
assert node.i(inp_idx).op == "DequantizeLinear", (
f"Input '{inp.name}' of node '{node.name}' is not quantized but should be!"
)
return True
@pytest.mark.parametrize("keep_io_types", [True, False])
@pytest.mark.parametrize("bias_add", [True, False])
def test_autocast_quantize_int8(tmp_path, keep_io_types, bias_add):
@@ -40,6 +40,17 @@ from modelopt.torch.export.quant_aware_conversion import (
BLOCK = 16
# Tiny Mixtral shaped to match the synthetic expert tensors built by ``_nvfp4_linear`` below.
_MIXTRAL_KWARGS = {
"hidden_size": 32,
"intermediate_size": 64,
"num_hidden_layers": 1,
"num_local_experts": 2,
"num_experts_per_tok": 2,
"vocab_size": 64,
"max_position_embeddings": 64,
}
def _nvfp4_linear(module: str, out: int, in_features: int) -> dict[str, torch.Tensor]:
"""Synthetic NVFP4 quantized-linear tensor group keyed under ``module``."""
@@ -176,8 +187,9 @@ def test_build_reverse_rules_from_mixtral_conversion_mapping_cpu():
a ModelOpt-expanded per-expert state dict (in-memory ``mlp.experts.<i>.*`` names)
must revert to the hub layout (``block_sparse_moe.experts.<i>.w{1,2,3}``).
"""
# Imports stay function-local: unit tests must import without transformers installed.
pytest.importorskip("transformers")
from transformers import MixtralConfig, MixtralForCausalLM
from _test_utils.torch.transformers_models import get_tiny_mixtral
try:
from transformers.conversion_mapping import get_checkpoint_conversion_mapping
@@ -186,18 +198,7 @@ def test_build_reverse_rules_from_mixtral_conversion_mapping_cpu():
if not get_checkpoint_conversion_mapping("mixtral"):
pytest.skip("transformers build has no mixtral conversion_mapping")
cfg = MixtralConfig(
hidden_size=32,
intermediate_size=64,
num_hidden_layers=1,
num_attention_heads=4,
num_key_value_heads=2,
num_local_experts=2,
num_experts_per_tok=2,
vocab_size=64,
max_position_embeddings=64,
)
model = MixtralForCausalLM(cfg)
model = get_tiny_mixtral(**_MIXTRAL_KWARGS)
p = "model.layers.0"
sd = {f"{p}.mlp.gate.weight": torch.randn(2, 32)}
@@ -354,22 +355,11 @@ def test_revert_quant_config_names_mapper():
deployment loader matched none of the excludes and loaded an excluded BF16 layer as
quantized. Uses Mixtral's real mapping (``mlp.experts`` <-> ``block_sparse_moe.experts``).
"""
# Import stays function-local: the helper needs transformers, which unit tests run without.
pytest.importorskip("transformers.core_model_loading")
from transformers import MixtralConfig, MixtralForCausalLM
from _test_utils.torch.transformers_models import get_tiny_mixtral
model = MixtralForCausalLM(
MixtralConfig(
hidden_size=32,
intermediate_size=64,
num_hidden_layers=1,
num_attention_heads=4,
num_key_value_heads=2,
num_local_experts=2,
num_experts_per_tok=2,
vocab_size=64,
max_position_embeddings=64,
)
)
model = get_tiny_mixtral(**_MIXTRAL_KWARGS)
mapper = build_reverse_name_mapper(model)
assert mapper is not None
@@ -32,7 +32,8 @@ import modelopt.torch.opt as mto
(AutoModelForCausalLM, "qwen3"),
],
)
def test_nested_model_save_restore(tmp_path, model_cls, teacher_model_type):
# Skipped on Windows - Flaky; root cause unknown; not critical
def test_nested_model_save_restore(skip_on_windows, tmp_path, model_cls, teacher_model_type):
tiny_llama_dir = create_tiny_llama_dir(tmp_path)
model_ref = model_cls.from_pretrained(tiny_llama_dir)
@@ -17,9 +17,8 @@ import pytest
import torch
import torch.nn as nn
import torch.nn.functional as F
from _test_utils.torch.quantization.attention import make_quant_attention
from _test_utils.torch.transformers_models import get_tiny_bert, get_tiny_llama, get_tiny_t5
from transformers import LlamaConfig
from transformers.models.llama.modeling_llama import LlamaAttention
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.plugins.huggingface import _QuantAttention
@@ -152,20 +151,9 @@ def test_kv_quant_bert():
assert output.end_logits is not None
def _make_quant_attention(hidden_size=128, num_q_heads=4, num_kv_heads=2):
config = LlamaConfig(
hidden_size=hidden_size,
num_attention_heads=num_q_heads,
num_key_value_heads=num_kv_heads,
)
quant_attention = _QuantAttention.convert(LlamaAttention(config, layer_idx=0))
quant_attention.config._attn_implementation = "sdpa"
return quant_attention
def test_p_qdq_mode_detection():
"""p_bmm_quantizer config maps to the right Triton softmax qdq mode."""
quant_attention = _make_quant_attention()
quant_attention = make_quant_attention()
sq = quant_attention.p_bmm_quantizer
# Default int8 quantizer: not a supported Triton qdq format
@@ -21,12 +21,12 @@ GPU-dependent tests (training forward, module forward) are in tests/gpu/.
import json
import logging
import os
from copy import deepcopy
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
import torch
from _test_utils.torch.speculative.dflash import get_dflash_config
from _test_utils.torch.transformers_models import (
get_tiny_llama,
tf_modelopt_state_and_output_tester,
@@ -36,7 +36,6 @@ from transformers import AutoModelForCausalLM
import modelopt.torch.opt as mto
import modelopt.torch.speculative as mtsp
import modelopt.torch.speculative.plugins.hf_dflash as hf_dflash
from modelopt.torch.speculative.config import DFLASH_DEFAULT_CFG
from modelopt.torch.speculative.plugins.hf_dflash import (
DFlashAttention,
DFlashModule,
@@ -52,32 +51,20 @@ NUM_DRAFT_LAYERS = 2
SEQ_LEN = 16 # must be multiple of BLOCK_SIZE
def _get_dflash_config(block_size=BLOCK_SIZE, num_layers=NUM_DRAFT_LAYERS):
"""Create a DFlash config for testing."""
config = deepcopy(DFLASH_DEFAULT_CFG["config"])
config["dflash_block_size"] = block_size
config["dflash_use_torch_compile"] = False
config["dflash_mask_token_id"] = 0 # use token 0 as mask for tiny model
config["dflash_architecture_config"] = {
"num_hidden_layers": num_layers,
}
return config
class TestDFlashConvert:
"""Test DFlash model conversion."""
def test_convert_creates_dflash_model(self):
"""Test that convert produces an HFDFlashModel."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
assert isinstance(model, HFDFlashModel)
def test_convert_creates_dflash_module(self):
"""Test that convert attaches a DFlashModule."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
assert hasattr(model, "dflash_module")
assert isinstance(model.dflash_module, DFlashModule)
@@ -85,7 +72,7 @@ class TestDFlashConvert:
def test_convert_freezes_base_model(self):
"""Test that base model parameters are frozen after convert."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
for name, param in model.named_parameters():
if "dflash_module" not in name:
@@ -94,7 +81,7 @@ class TestDFlashConvert:
def test_convert_dflash_module_trainable(self):
"""Test that DFlash module parameters are trainable after convert."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
dflash_params = [(n, p) for n, p in model.named_parameters() if "dflash_module" in n]
assert len(dflash_params) > 0
@@ -104,7 +91,7 @@ class TestDFlashConvert:
def test_convert_sets_target_layer_ids(self):
"""Test that target layer IDs are set correctly."""
model = get_tiny_llama(num_hidden_layers=8)
config = _get_dflash_config(num_layers=3)
config = get_dflash_config(num_layers=3)
mtsp.convert(model, [("dflash", config)])
assert hasattr(model, "target_layer_ids")
assert len(model.target_layer_ids) == 3
@@ -114,7 +101,7 @@ class TestDFlashConvert:
def test_convert_sets_mask_token_id(self):
"""Test that mask_token_id is set from config."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
assert hasattr(model, "mask_token_id")
assert model.mask_token_id == 0
@@ -283,7 +270,7 @@ def test_multimodal_forward_kwargs_exclude_non_model_inputs():
def test_eval_does_not_precompute_qwen3_vl_position_ids(monkeypatch):
"""Evaluation delegates mRoPE construction to the base model and its cache."""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash_config())])
mtsp.convert(model, [("dflash", get_dflash_config())])
precompute_position_ids = MagicMock()
monkeypatch.setattr(model, "_qwen3_vl_position_ids", precompute_position_ids)
@@ -395,26 +382,26 @@ class TestDPaceWeights:
def test_default_objective_is_dpace(self):
"""D-PACE is the default (alpha=0.5); an explicit alpha override is wired through."""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash_config())])
mtsp.convert(model, [("dflash", get_dflash_config())])
assert model.dflash_loss_objective == "dpace"
assert model.dflash_dpace_alpha == 0.5
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
config["dflash_dpace_alpha"] = 0.3
mtsp.convert(model, [("dflash", config)])
assert model.dflash_dpace_alpha == 0.3
def test_convert_rejects_bad_objective(self):
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
config["dflash_loss_objective"] = "nope"
with pytest.raises(ValueError, match="dflash_loss_objective"):
mtsp.convert(model, [("dflash", config)])
def test_convert_rejects_degenerate_alpha(self):
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
config["dflash_loss_objective"] = "dpace"
config["dflash_dpace_alpha"] = 0.0
with pytest.raises(ValueError, match="dflash_dpace_alpha"):
@@ -423,7 +410,7 @@ class TestDPaceWeights:
def test_convert_dpace_with_decay_factor_warns(self, caplog):
"""dpace + a non-zero decay factor converts but warns that decay is ignored."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
config["dflash_loss_objective"] = "dpace"
config["dflash_loss_decay_factor"] = 4.0
with caplog.at_level(logging.WARNING):
@@ -447,7 +434,7 @@ class TestDPaceLossIntegration:
def _converted_model(self, objective, **overrides):
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
config["dflash_loss_objective"] = objective
config.update(overrides)
mtsp.convert(model, [("dflash", config)])
@@ -485,7 +472,7 @@ class TestDFlashSaveRestore:
"""Test round-trip save/load preserves modelopt state and outputs."""
mto.enable_huggingface_checkpointing()
model_ref = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model_ref, [("dflash", config)])
model_ref.save_pretrained(tmp_path / "modelopt_model")
@@ -506,14 +493,14 @@ class TestDFlashLazyRotaryEmb:
def test_rotary_emb_not_created_in_init(self):
"""rotary_emb should not exist after convert (before forward)."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
assert not hasattr(model.dflash_module, "rotary_emb")
def test_rotary_emb_created_on_forward(self):
"""rotary_emb should be created on first forward call."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
dflash_mod = model.dflash_module
@@ -595,7 +582,7 @@ class TestDFlashSwaMask:
def test_window_masks_context_beyond_window(self):
"""Context beyond the window (relative to each query's real position) is masked out."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config(block_size=4)
config = get_dflash_config(block_size=4)
window = 6
config["dflash_swa_window_size"] = window
mtsp.convert(model, [("dflash", config)])
@@ -629,7 +616,7 @@ class TestDFlashSwaMask:
def test_window_is_subset_of_full(self):
"""The windowed mask attends to a subset of what the full-attention mask attends to."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config(block_size=4)
config = get_dflash_config(block_size=4)
config["dflash_swa_window_size"] = 6
mtsp.convert(model, [("dflash", config)])
@@ -652,7 +639,7 @@ class TestDFlashSwaMask:
def test_window_smaller_than_block_rejected(self):
"""A window smaller than the block size is rejected at config validation."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config(block_size=4)
config = get_dflash_config(block_size=4)
config["dflash_swa_window_size"] = 2 # < block_size
with pytest.raises(ValueError, match="dflash_swa_window_size"):
mtsp.convert(model, [("dflash", config)])
@@ -748,7 +735,7 @@ class TestDFlashExporter:
"""Test that export produces model.safetensors and config.json."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
exporter = model.get_exporter()
@@ -763,7 +750,7 @@ class TestDFlashExporter:
from safetensors.torch import load_file
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
exporter = model.get_exporter()
@@ -779,7 +766,7 @@ class TestDFlashExporter:
def test_export_config_fields(self, tmp_path):
"""Exported config.json should have required DFlash fields."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
exporter = model.get_exporter()
@@ -806,7 +793,7 @@ class TestDFlashExporter:
def test_export_swa_fields(self, tmp_path):
"""With dflash_swa_window_size set, exported config carries vLLM's SWA fields."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
config["dflash_swa_window_size"] = 256
mtsp.convert(model, [("dflash", config)])
@@ -832,7 +819,7 @@ class TestDFlashExporter:
from safetensors.torch import load_file
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash_config()
config = get_dflash_config()
mtsp.convert(model, [("dflash", config)])
exporter = model.get_exporter()
@@ -15,33 +15,20 @@
"""CPU unit tests for DFlash offline training support."""
from copy import deepcopy
from _test_utils.torch.speculative.dflash import get_dflash_config
from _test_utils.torch.transformers_models import get_tiny_llama
import modelopt.torch.speculative as mtsp
from modelopt.recipe.config import ModelOptDFlashRecipe
from modelopt.torch.speculative.config import DFLASH_DEFAULT_CFG
NUM_BASE_LAYERS = 4
NUM_DRAFT_LAYERS = 2
def _get_dflash_config(offline=False):
"""Build a minimal DFlash config dict for mtsp.convert."""
config = deepcopy(DFLASH_DEFAULT_CFG["config"])
config["dflash_offline"] = offline
config["dflash_block_size"] = 4
config["dflash_use_torch_compile"] = False
config["dflash_mask_token_id"] = 0
config["dflash_architecture_config"] = {"num_hidden_layers": NUM_DRAFT_LAYERS}
return config
def test_convert_online_keeps_base_layers():
"""Online DFlash (default) keeps the base model layers intact."""
model = get_tiny_llama(num_hidden_layers=NUM_BASE_LAYERS)
mtsp.convert(model, [("dflash", _get_dflash_config(offline=False))])
mtsp.convert(model, [("dflash", get_dflash_config(offline=False))])
assert model.dflash_offline is False
assert "layers" in model._base_model._modules
@@ -53,7 +40,7 @@ def test_convert_offline_deletes_base_layers():
model = get_tiny_llama(num_hidden_layers=NUM_BASE_LAYERS)
# num_orig_hidden_layers records the pre-deletion layer count; users set it before convert.
model.config.num_orig_hidden_layers = NUM_BASE_LAYERS
mtsp.convert(model, [("dflash", _get_dflash_config(offline=True))])
mtsp.convert(model, [("dflash", get_dflash_config(offline=True))])
assert model.dflash_offline is True
assert "layers" not in model._base_model._modules
@@ -64,7 +51,7 @@ def test_convert_offline_target_layer_ids_from_orig():
num_orig = 8
model = get_tiny_llama(num_hidden_layers=NUM_BASE_LAYERS)
model.config.num_orig_hidden_layers = num_orig
mtsp.convert(model, [("dflash", _get_dflash_config(offline=True))])
mtsp.convert(model, [("dflash", get_dflash_config(offline=True))])
assert len(model.target_layer_ids) == NUM_DRAFT_LAYERS
# With num_orig=8, build_target_layer_ids(8, 2) spans beyond the 4 live base layers —