mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
ModelOpt Framework, Recipe Lib, converting subset of existing recipes 1/N (#1000)
### What does this PR do?
1. start a new config system using yaml/yml files.
2. add a new top level package: modelopt_recipes
I want it to be a top level package so we can make it clear that the
modelopt package holds the code, this new package holds recipes
3. implement some of the existing quantization recipes using the new
config system as model agnostic general recipes, but not actually in
use. these recipes sit inside modelopt_recipes/general/ptq/...
4. make sure the configs from the new config system match the exisiting
configs
5. extend the hf_ptq script to enable recipe based PTQ
8. testted hf_ptq using both builtin and extenal config file. example
script:
### Usage
```bash
python examples/llm_ptq/hf_ptq.py \
--model Qwen/Qwen3-8B \
--recipe general/ptq/fp8_default-fp8_kv \
...
```
### Testing
```bash
python examples/llm_ptq/hf_ptq.py \
--model Qwen/Qwen3-8B \
--recipe general/ptq/fp8_default-fp8_kv \
--export_path=fp8_default-fp8_kv \
--calib_size=16 \
--batch_size=0 \
--trust_remote_code \
--export_fmt=hf
```
### Before your PR is "*Ready for review*"
Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).
Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).
- Is this change backward compatible?: ✅ / ❌ / N/A <!--- If ❌, explain
why. -->
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ / ❌ / N/A
<!--- Mandatory -->
- Did you write any new necessary tests?: ✅ / ❌ / N/A <!--- Mandatory
for new features or examples. -->
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ / ❌ / N/A <!--- Only for new features, API changes, critical bug fixes
or backward incompatible changes. -->
### Additional Information
<!-- E.g. related issue. -->
<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit
* **New Features**
* Recipe-driven PTQ workflows via YAML recipes and new recipe loader;
CLI gains a --recipe option and --pyt_ckpt_path renamed to --model.
* Many new PTQ recipe and config presets (FP8, INT4/INT8, NVFP4, MXFPx,
KV-cache variants) and improved runtime config loading/merging.
* **Documentation**
* Added READMEs describing recipe/config layout.
* **Tests**
* New unit tests covering config loading, inheritance and recipe
loading.
* **Chores**
* Added YAML/OmegaConf runtime support and packaging of recipe YAMLs.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
---------
Signed-off-by: Shengliang Xu <shengliangx@nvidia.com>
This commit is contained in:
@@ -45,7 +45,6 @@ try:
|
||||
except ImportError:
|
||||
snapshot_download = None
|
||||
|
||||
import modelopt.torch.quantization as mtq
|
||||
from modelopt.torch.utils.image_processor import BaseImageProcessor, MllamaImageProcessor
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -199,22 +198,13 @@ def create_vlm_calibration_loop(full_model, calib_dataloader):
|
||||
|
||||
def build_quant_cfg(
|
||||
qformat,
|
||||
kv_cache_qformat,
|
||||
quant_cfg,
|
||||
awq_block_size,
|
||||
model_type,
|
||||
quant_cfg_choices,
|
||||
kv_quant_cfg_choices,
|
||||
moe_calib_experts_ratio: float | None = None,
|
||||
) -> dict[str, Any]:
|
||||
quant_cfg = {}
|
||||
assert qformat in quant_cfg_choices, (
|
||||
f"Unsupported quantization format: {qformat} with {kv_cache_qformat} KV cache"
|
||||
)
|
||||
|
||||
quant_cfg = quant_cfg_choices[qformat]
|
||||
|
||||
if "awq" in qformat:
|
||||
quant_cfg = copy.deepcopy(quant_cfg_choices[qformat])
|
||||
quant_cfg = copy.deepcopy(quant_cfg)
|
||||
if "awq" in str(quant_cfg.get("algorithm")):
|
||||
weight_quantizer = quant_cfg["quant_cfg"]["*weight_quantizer"]
|
||||
if isinstance(weight_quantizer, list):
|
||||
weight_quantizer = weight_quantizer[0]
|
||||
@@ -226,16 +216,6 @@ def build_quant_cfg(
|
||||
if qformat == "w4a8_awq" and model_type in ["gemma", "mpt"]:
|
||||
quant_cfg["algorithm"] = {"method": "awq_lite", "alpha_step": 1}
|
||||
|
||||
enable_quant_kv_cache = kv_cache_qformat != "none"
|
||||
print(f"{'Enable' if enable_quant_kv_cache else 'Disable'} KV cache quantization")
|
||||
|
||||
# Check if any bmm_quantizer is in the quant_cfg. If so, we need to enable the bmm_quantizer.
|
||||
if enable_quant_kv_cache:
|
||||
quant_cfg = mtq.update_quant_cfg_with_kv_cache_quant(
|
||||
quant_cfg,
|
||||
getattr(mtq, kv_quant_cfg_choices[kv_cache_qformat])["quant_cfg"],
|
||||
)
|
||||
|
||||
if moe_calib_experts_ratio:
|
||||
assert 0 < moe_calib_experts_ratio <= 1, "moe_calib_experts_ratio must be between 0 and 1"
|
||||
if isinstance(quant_cfg["algorithm"], str):
|
||||
|
||||
+50
-38
@@ -50,6 +50,7 @@ from transformers import (
|
||||
import modelopt.torch.opt as mto
|
||||
import modelopt.torch.quantization as mtq
|
||||
import modelopt.torch.sparsity as mts
|
||||
from modelopt.recipe import ModelOptPTQRecipe, load_recipe
|
||||
from modelopt.torch.export import (
|
||||
export_hf_checkpoint,
|
||||
export_speculative_decoding,
|
||||
@@ -262,7 +263,7 @@ def auto_quantize(
|
||||
assert qformat_list, "No quantization formats provided"
|
||||
# Check if all provided quantization formats are supported
|
||||
assert all(
|
||||
args.qformat
|
||||
qformat
|
||||
in [
|
||||
"fp8",
|
||||
"int8_sq",
|
||||
@@ -277,7 +278,7 @@ def auto_quantize(
|
||||
"nvfp4_omlp_only",
|
||||
"mxfp8",
|
||||
]
|
||||
for args.qformat in qformat_list
|
||||
for qformat in qformat_list
|
||||
), "One or more quantization formats provided are not supported for unified checkpoint export"
|
||||
|
||||
def loss_func(output, data):
|
||||
@@ -548,9 +549,6 @@ def mono_quantize(
|
||||
print("Quantization will only be applied to the decoder (text generation) component")
|
||||
|
||||
if not model_is_already_quantized or calibration_only:
|
||||
if model_type == "gptoss" and args.qformat == "nvfp4_mlp_only":
|
||||
print("Applying nvfp4 quantization (MoE only) for gpt-oss")
|
||||
|
||||
# quantize the model
|
||||
|
||||
use_calibration = need_calibration(quant_cfg)
|
||||
@@ -746,8 +744,6 @@ def pre_quantize(
|
||||
)
|
||||
else:
|
||||
generated_ids_before_ptq = full_model.generate(preview_input_ids, max_new_tokens=100)
|
||||
if model_type == "gptoss" and args.qformat == "nvfp4_mlp_only":
|
||||
print("Applying nvfp4 quantization (MoE only) for gpt-oss")
|
||||
|
||||
return preview_input_ids, generated_ids_before_ptq
|
||||
|
||||
@@ -923,38 +919,42 @@ def quantize_main(
|
||||
|
||||
else:
|
||||
# mono quantization
|
||||
assert len(args.qformat.split(",")) == 1, (
|
||||
"Plain quantization supports only one quantization format."
|
||||
)
|
||||
|
||||
assert (
|
||||
args.qformat
|
||||
in [
|
||||
"int8_wo",
|
||||
"int4_awq",
|
||||
"fp8",
|
||||
"nvfp4",
|
||||
"nvfp4_awq",
|
||||
"nvfp4_mse",
|
||||
"w4a8_awq",
|
||||
"fp8_pb_wo",
|
||||
"w4a8_mxfp4_fp8",
|
||||
"nvfp4_mlp_only",
|
||||
"nvfp4_omlp_only",
|
||||
"mxfp8",
|
||||
]
|
||||
or args.kv_cache_qformat in KV_QUANT_CFG_CHOICES
|
||||
), f"Plain quantization format {args.qformat} not supported for HF export path"
|
||||
if args.recipe is not None:
|
||||
print(f"Use recipe {args.recipe} for quantization")
|
||||
recipe = load_recipe(args.recipe)
|
||||
assert isinstance(recipe, ModelOptPTQRecipe), (
|
||||
f"Expected PTQ recipe, but got {type(recipe).__name__} from {args.recipe}"
|
||||
)
|
||||
quant_cfg = recipe.ptq_cfg
|
||||
|
||||
quant_cfg = build_quant_cfg(
|
||||
args.qformat,
|
||||
args.kv_cache_qformat,
|
||||
args.awq_block_size,
|
||||
model_type,
|
||||
QUANT_CFG_CHOICES,
|
||||
KV_QUANT_CFG_CHOICES,
|
||||
args.moe_calib_experts_ratio,
|
||||
)
|
||||
else:
|
||||
assert len(args.qformat.split(",")) == 1, (
|
||||
"Plain quantization supports only one quantization format."
|
||||
)
|
||||
|
||||
assert args.qformat in QUANT_CFG_CHOICES, (
|
||||
f"Unsupported quantization format: {args.qformat}, choices are: {list(QUANT_CFG_CHOICES.keys())}"
|
||||
)
|
||||
quant_cfg = QUANT_CFG_CHOICES[args.qformat]
|
||||
|
||||
quant_cfg = build_quant_cfg(
|
||||
args.qformat,
|
||||
quant_cfg,
|
||||
args.awq_block_size,
|
||||
model_type,
|
||||
args.moe_calib_experts_ratio,
|
||||
)
|
||||
|
||||
enable_quant_kv_cache = args.kv_cache_qformat != "none"
|
||||
print(f"{'Enable' if enable_quant_kv_cache else 'Disable'} KV cache quantization")
|
||||
|
||||
# Check if any bmm_quantizer is in the quant_cfg. If so, we need to enable the bmm_quantizer.
|
||||
if enable_quant_kv_cache:
|
||||
quant_cfg = mtq.update_quant_cfg_with_kv_cache_quant(
|
||||
quant_cfg,
|
||||
getattr(mtq, KV_QUANT_CFG_CHOICES[args.kv_cache_qformat])["quant_cfg"],
|
||||
)
|
||||
|
||||
# Exclude MTP layers from quantization if detected (e.g., GLM-4.7's layer 92)
|
||||
# These layers are typically speculative decoding layers that should be exported as-is
|
||||
@@ -1013,9 +1013,21 @@ def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--pyt_ckpt_path",
|
||||
help="Specify where the PyTorch checkpoint path is",
|
||||
"--model",
|
||||
help=(
|
||||
"Model name or path to the PyTorch checkpoint to be quantized. "
|
||||
"Can be a local path or a Huggingface model name."
|
||||
),
|
||||
required=True,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--recipe",
|
||||
help=(
|
||||
"PTQ recipe YAML file or name without suffix (e.g. general/ptq/nvfp4_default-fp8_kv)."
|
||||
),
|
||||
default=None,
|
||||
)
|
||||
|
||||
parser.add_argument("--device", default="cuda")
|
||||
parser.add_argument(
|
||||
"--qformat",
|
||||
|
||||
@@ -327,16 +327,25 @@ def main(args):
|
||||
trust_remote_code=args.trust_remote_code,
|
||||
)
|
||||
|
||||
# Build quantization config
|
||||
quant_cfg = QUANT_CFG_CHOICES[args.qformat]
|
||||
|
||||
quant_cfg = build_quant_cfg(
|
||||
args.qformat,
|
||||
args.kv_cache_qformat,
|
||||
quant_cfg,
|
||||
args.awq_block_size,
|
||||
model_type,
|
||||
QUANT_CFG_CHOICES,
|
||||
KV_QUANT_CFG_CHOICES,
|
||||
)
|
||||
|
||||
enable_quant_kv_cache = args.kv_cache_qformat != "none"
|
||||
print(f"{'Enable' if enable_quant_kv_cache else 'Disable'} KV cache quantization")
|
||||
|
||||
# Check if any bmm_quantizer is in the quant_cfg. If so, we need to enable the bmm_quantizer.
|
||||
if enable_quant_kv_cache:
|
||||
quant_cfg = mtq.update_quant_cfg_with_kv_cache_quant(
|
||||
quant_cfg,
|
||||
getattr(mtq, KV_QUANT_CFG_CHOICES[args.kv_cache_qformat])["quant_cfg"],
|
||||
)
|
||||
|
||||
# Quantize the model
|
||||
if accelerator.is_main_process:
|
||||
print("Starting quantization...")
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
# 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.
|
||||
|
||||
"""Module for the ModelOpt recipe lib.
|
||||
|
||||
``modelopt.recipe`` contains tooling to:
|
||||
|
||||
* load and store model optimization recipes
|
||||
* (TODO) utilities to manipulate the recipes, such as merging multiple recipes together, or
|
||||
overriding some fields in a recipe with user-provided values.
|
||||
|
||||
"""
|
||||
|
||||
from .config import *
|
||||
from .loader import *
|
||||
@@ -0,0 +1,113 @@
|
||||
# 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.
|
||||
|
||||
"""YAML config loading utilities.
|
||||
|
||||
This module is intentionally free of ``modelopt.torch`` imports so that
|
||||
``modelopt.torch.quantization.config`` can import :func:`load_config` without
|
||||
triggering a circular import through ``modelopt.recipe.loader``.
|
||||
"""
|
||||
|
||||
from importlib.resources import files
|
||||
|
||||
try:
|
||||
from importlib.resources.abc import Traversable
|
||||
except ImportError: # Python < 3.11
|
||||
from importlib.abc import Traversable
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
|
||||
# Root to all built-in recipes. Users can create own recipes.
|
||||
BUILTIN_RECIPES_LIB = files("modelopt_recipes")
|
||||
|
||||
_EXMY_RE = re.compile(r"^[Ee](\d+)[Mm](\d+)$")
|
||||
_EXMY_KEYS = frozenset({"num_bits", "scale_bits"})
|
||||
|
||||
|
||||
def _parse_exmy_num_bits(obj: Any) -> Any:
|
||||
"""Recursively convert ``ExMy`` strings in ``num_bits`` / ``scale_bits`` to ``(x, y)`` tuples."""
|
||||
if isinstance(obj, dict):
|
||||
return {
|
||||
k: (
|
||||
_parse_exmy(v)
|
||||
if k in _EXMY_KEYS and isinstance(v, str)
|
||||
else _parse_exmy_num_bits(v)
|
||||
)
|
||||
for k, v in obj.items()
|
||||
}
|
||||
if isinstance(obj, list):
|
||||
return [_parse_exmy_num_bits(item) for item in obj]
|
||||
return obj
|
||||
|
||||
|
||||
def _parse_exmy(s: str) -> tuple[int, int] | str:
|
||||
m = _EXMY_RE.match(s)
|
||||
if m:
|
||||
return (int(m.group(1)), int(m.group(2)))
|
||||
return s
|
||||
|
||||
|
||||
def load_config(config_file: str | Path | Traversable) -> dict[str, Any]:
|
||||
"""Load a config yaml.
|
||||
|
||||
config_file: Path to a config yaml file. The path suffix can be omitted.
|
||||
"""
|
||||
paths_to_check: list[Path | Traversable] = []
|
||||
if isinstance(config_file, str):
|
||||
if not config_file.endswith(".yml") and not config_file.endswith(".yaml"):
|
||||
paths_to_check.append(Path(f"{config_file}.yml"))
|
||||
paths_to_check.append(Path(f"{config_file}.yaml"))
|
||||
paths_to_check.append(BUILTIN_RECIPES_LIB.joinpath(f"{config_file}.yml"))
|
||||
paths_to_check.append(BUILTIN_RECIPES_LIB.joinpath(f"{config_file}.yaml"))
|
||||
else:
|
||||
paths_to_check.append(Path(config_file))
|
||||
paths_to_check.append(BUILTIN_RECIPES_LIB.joinpath(config_file))
|
||||
elif isinstance(config_file, Path):
|
||||
if config_file.suffix in (".yml", ".yaml"):
|
||||
paths_to_check.append(config_file)
|
||||
if not config_file.is_absolute():
|
||||
paths_to_check.append(BUILTIN_RECIPES_LIB.joinpath(str(config_file)))
|
||||
else:
|
||||
paths_to_check.append(Path(f"{config_file}.yml"))
|
||||
paths_to_check.append(Path(f"{config_file}.yaml"))
|
||||
if not config_file.is_absolute():
|
||||
paths_to_check.append(BUILTIN_RECIPES_LIB.joinpath(f"{config_file}.yml"))
|
||||
paths_to_check.append(BUILTIN_RECIPES_LIB.joinpath(f"{config_file}.yaml"))
|
||||
elif isinstance(config_file, Traversable):
|
||||
paths_to_check.append(config_file)
|
||||
else:
|
||||
raise ValueError(f"Invalid config file of {config_file}")
|
||||
|
||||
config_path = None
|
||||
for path in paths_to_check:
|
||||
if path.is_file():
|
||||
config_path = path
|
||||
break
|
||||
if not config_path:
|
||||
raise ValueError(
|
||||
f"Cannot find config file of {config_file}, paths checked: {paths_to_check}"
|
||||
)
|
||||
|
||||
_raw = yaml.safe_load(config_path.read_text(encoding="utf-8"))
|
||||
if _raw is None:
|
||||
return {}
|
||||
if not isinstance(_raw, dict):
|
||||
raise ValueError(
|
||||
f"Config file {config_path} must contain a YAML mapping, got {type(_raw).__name__}"
|
||||
)
|
||||
return _parse_exmy_num_bits(_raw)
|
||||
@@ -0,0 +1,74 @@
|
||||
# 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.
|
||||
|
||||
"""ModelOpt's pydantic BaseModel for recipes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
from pydantic import field_validator
|
||||
|
||||
from modelopt.torch.opt.config import ModeloptBaseConfig, ModeloptField
|
||||
|
||||
|
||||
class RecipeType(str, Enum):
|
||||
"""List of recipe types."""
|
||||
|
||||
PTQ = "ptq"
|
||||
# QAT = "qat" # Not implemented yet, will be added in the future.
|
||||
|
||||
|
||||
class ModelOptRecipeBase(ModeloptBaseConfig):
|
||||
"""Base configuration class for model optimization recipes.
|
||||
|
||||
If a layer name matches ``"*output_layer*"``, the attributes will be replaced with ``{"enable": False}``.
|
||||
"""
|
||||
|
||||
recipe_type: RecipeType = ModeloptField(
|
||||
default=RecipeType.PTQ,
|
||||
title="type of the recipe",
|
||||
description="The type of the recipe.",
|
||||
validate_default=True,
|
||||
)
|
||||
|
||||
description: str = ModeloptField(
|
||||
default="Model optimization recipe.",
|
||||
title="Description",
|
||||
description="A brief description of the model optimization recipe.",
|
||||
validate_default=False,
|
||||
)
|
||||
|
||||
@field_validator("recipe_type")
|
||||
@classmethod
|
||||
def validate_recipe_type(cls, v):
|
||||
"""Validate recipe type."""
|
||||
if v not in RecipeType:
|
||||
raise ValueError(
|
||||
f"Unsupported recipe type: {v}. Only {list(RecipeType)} are currently supported."
|
||||
)
|
||||
return v
|
||||
|
||||
|
||||
class ModelOptPTQRecipe(ModelOptRecipeBase):
|
||||
"""Our config class for PTQ recipes."""
|
||||
|
||||
ptq_cfg: dict[str, Any] = ModeloptField(
|
||||
default={},
|
||||
title="PTQ config",
|
||||
description="PTQ config containing quant_cfg and algorithm.",
|
||||
validate_default=True,
|
||||
)
|
||||
@@ -0,0 +1,141 @@
|
||||
# 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.
|
||||
|
||||
"""Recipe loading utilities."""
|
||||
|
||||
try:
|
||||
from importlib.resources.abc import Traversable
|
||||
except ImportError: # Python < 3.11
|
||||
from importlib.abc import Traversable
|
||||
from pathlib import Path
|
||||
|
||||
from ._config_loader import BUILTIN_RECIPES_LIB, load_config
|
||||
from .config import ModelOptPTQRecipe, ModelOptRecipeBase, RecipeType
|
||||
|
||||
__all__ = ["load_config", "load_recipe"]
|
||||
|
||||
|
||||
def _resolve_recipe_path(recipe_path: str | Path | Traversable) -> Path | Traversable:
|
||||
"""Resolve a recipe path, checking the built-in library first then the filesystem.
|
||||
|
||||
Returns the resolved path (file or directory).
|
||||
"""
|
||||
if isinstance(recipe_path, (str, Path)) and not (
|
||||
isinstance(recipe_path, Path) and recipe_path.is_absolute()
|
||||
):
|
||||
rp_str = str(recipe_path)
|
||||
suffixes = [""] if rp_str.endswith((".yml", ".yaml")) else ["", ".yml", ".yaml"]
|
||||
for suffix in suffixes:
|
||||
candidate = BUILTIN_RECIPES_LIB.joinpath(rp_str + suffix)
|
||||
if candidate.is_file() or candidate.is_dir():
|
||||
return candidate
|
||||
for suffix in suffixes:
|
||||
fs_candidate = Path(rp_str + suffix)
|
||||
if fs_candidate.is_file() or fs_candidate.is_dir():
|
||||
return fs_candidate
|
||||
return Path(rp_str)
|
||||
return recipe_path
|
||||
|
||||
|
||||
def load_recipe(recipe_path: str | Path | Traversable) -> ModelOptRecipeBase:
|
||||
"""Load a recipe from a YAML file or directory.
|
||||
|
||||
``recipe_path`` can be:
|
||||
|
||||
* A ``.yml`` / ``.yaml`` file with ``metadata`` and ``ptq_cfg`` sections.
|
||||
The suffix may be omitted and will be probed automatically.
|
||||
* A directory containing ``recipe.yml`` (metadata) and ``ptq_cfg.yml``.
|
||||
|
||||
The path may be relative to the built-in recipes library or an absolute /
|
||||
relative filesystem path.
|
||||
"""
|
||||
resolved = _resolve_recipe_path(recipe_path)
|
||||
|
||||
_builtin_prefix = str(BUILTIN_RECIPES_LIB)
|
||||
_resolved_str = str(resolved)
|
||||
if _resolved_str.startswith(_builtin_prefix):
|
||||
_display = "<builtin>/" + _resolved_str[len(_builtin_prefix) :].lstrip("/\\")
|
||||
else:
|
||||
_display = _resolved_str
|
||||
print(f"[load_recipe] loading: {_display}")
|
||||
|
||||
if resolved.is_file():
|
||||
return _load_recipe_from_file(resolved)
|
||||
|
||||
if resolved.is_dir():
|
||||
return _load_recipe_from_dir(resolved)
|
||||
|
||||
raise ValueError(f"Recipe path {recipe_path!r} is not a valid YAML file or directory.")
|
||||
|
||||
|
||||
def _load_recipe_from_file(recipe_file: Path | Traversable) -> ModelOptRecipeBase:
|
||||
"""Load a recipe from a YAML file.
|
||||
|
||||
The file must contain a ``metadata`` section with at least ``recipe_type``,
|
||||
plus a ``quant_cfg`` mapping and an optional ``algorithm`` for PTQ recipes.
|
||||
"""
|
||||
data = load_config(recipe_file)
|
||||
|
||||
metadata = data.get("metadata", {})
|
||||
recipe_type = metadata.get("recipe_type")
|
||||
if recipe_type is None:
|
||||
raise ValueError(f"Recipe file {recipe_file} must contain a 'metadata.recipe_type' field.")
|
||||
|
||||
if recipe_type == RecipeType.PTQ:
|
||||
if "ptq_cfg" not in data:
|
||||
raise ValueError(f"PTQ recipe file {recipe_file} must contain 'ptq_cfg'.")
|
||||
return ModelOptPTQRecipe(
|
||||
recipe_type=RecipeType.PTQ,
|
||||
description=metadata.get("description", "PTQ recipe."),
|
||||
ptq_cfg=data["ptq_cfg"],
|
||||
)
|
||||
raise ValueError(f"Unsupported recipe type: {recipe_type!r}")
|
||||
|
||||
|
||||
def _load_recipe_from_dir(recipe_dir: Path | Traversable) -> ModelOptRecipeBase:
|
||||
"""Load a recipe from a directory containing ``recipe.yml`` and ``ptq_cfg.yml``."""
|
||||
recipe_file = None
|
||||
for name in ("recipe.yml", "recipe.yaml"):
|
||||
candidate = recipe_dir.joinpath(name)
|
||||
if candidate.is_file():
|
||||
recipe_file = candidate
|
||||
break
|
||||
if recipe_file is None:
|
||||
raise ValueError(
|
||||
f"Cannot find a recipe descriptor in {recipe_dir}. Looked for: recipe.yml, recipe.yaml"
|
||||
)
|
||||
|
||||
metadata = load_config(recipe_file).get("metadata", {})
|
||||
recipe_type = metadata.get("recipe_type")
|
||||
if recipe_type is None:
|
||||
raise ValueError(f"Recipe file {recipe_file} must contain a 'metadata.recipe_type' field.")
|
||||
|
||||
if recipe_type == RecipeType.PTQ:
|
||||
ptq_cfg_file = None
|
||||
for name in ("ptq_cfg.yml", "ptq_cfg.yaml"):
|
||||
candidate = recipe_dir.joinpath(name)
|
||||
if candidate.is_file():
|
||||
ptq_cfg_file = candidate
|
||||
break
|
||||
if ptq_cfg_file is None:
|
||||
raise ValueError(
|
||||
f"Cannot find ptq_cfg in {recipe_dir}. Looked for: ptq_cfg.yml, ptq_cfg.yaml"
|
||||
)
|
||||
return ModelOptPTQRecipe(
|
||||
recipe_type=RecipeType.PTQ,
|
||||
description=metadata.get("description", "PTQ recipe."),
|
||||
ptq_cfg=load_config(ptq_cfg_file),
|
||||
)
|
||||
raise ValueError(f"Unsupported recipe type: {recipe_type!r}")
|
||||
@@ -20,7 +20,6 @@ import json
|
||||
from collections.abc import Callable, ItemsView, Iterator, KeysView, ValuesView
|
||||
from typing import Any, TypeAlias
|
||||
|
||||
import pydantic
|
||||
from pydantic import (
|
||||
BaseModel,
|
||||
Field,
|
||||
@@ -30,6 +29,7 @@ from pydantic import (
|
||||
field_validator,
|
||||
model_validator,
|
||||
)
|
||||
from pydantic import ConfigDict as PyDanticConfigDict
|
||||
from pydantic_core import PydanticUndefined
|
||||
|
||||
# A simple type alias for a config dictionary that is used as input to initialize a ModeloptBaseConfig.
|
||||
@@ -63,7 +63,7 @@ class ModeloptBaseConfig(BaseModel):
|
||||
and properties for easier access and manipulation of the configuration.
|
||||
"""
|
||||
|
||||
model_config = pydantic.ConfigDict(extra="forbid", validate_assignment=True)
|
||||
model_config = PyDanticConfigDict(extra="forbid", validate_assignment=True)
|
||||
|
||||
def model_dump(self, **kwargs):
|
||||
"""Dump the config to a dictionary with aliases and no warnings by default."""
|
||||
@@ -214,7 +214,7 @@ class ModeloptBaseRuleConfig(ModeloptBaseConfig):
|
||||
and properties for easier access and manipulation of the configuration.
|
||||
"""
|
||||
|
||||
model_config = pydantic.ConfigDict(extra="allow")
|
||||
model_config = PyDanticConfigDict(extra="allow")
|
||||
|
||||
@classmethod
|
||||
def __init_subclass__(cls, *args, registry, **kwargs):
|
||||
|
||||
@@ -658,6 +658,8 @@ NVFP4_OMLP_ONLY_CFG = {
|
||||
"algorithm": "max",
|
||||
}
|
||||
|
||||
# DO NOT ADD NEW CONFIGS HERE. If you want to add a new general recipe, add it to
|
||||
# modelopt_recipes/general/ptq/ as a yaml file
|
||||
choices: set[str] = {
|
||||
"FP8_2D_BLOCKWISE_WEIGHT_ONLY_CFG",
|
||||
"FP8_AFFINE_KV_CFG",
|
||||
|
||||
@@ -15,6 +15,7 @@
|
||||
|
||||
"""Quantization utilities."""
|
||||
|
||||
import copy
|
||||
from collections import namedtuple
|
||||
from contextlib import ExitStack, contextmanager, nullcontext
|
||||
from typing import TYPE_CHECKING, Any
|
||||
@@ -826,7 +827,8 @@ def update_quant_cfg_with_kv_cache_quant(
|
||||
) -> dict[str, Any]:
|
||||
"""Update the quant_cfg with the kv cache quant_cfg."""
|
||||
# If quant_cfg["quant_cfg"] is None, it corresponds to only kv cache quantization case
|
||||
quant_cfg["quant_cfg"] = quant_cfg.get("quant_cfg", {"default": {"enable": False}})
|
||||
quant_cfg = copy.deepcopy(quant_cfg)
|
||||
quant_cfg["quant_cfg"] = quant_cfg.get("quant_cfg") or {"default": {"enable": False}}
|
||||
quant_cfg["quant_cfg"].update(kv_cache_quant_cfg)
|
||||
|
||||
# Set default algorithm for kv cache quantization if not provided.
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
# 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.
|
||||
|
||||
"""Nvidia Model Optimizer Recipes Repo.
|
||||
|
||||
We need to make it a python package to be able to load the recipe resources programmatically.
|
||||
"""
|
||||
|
||||
import modelopt
|
||||
|
||||
__version__ = modelopt.__version__
|
||||
@@ -0,0 +1,3 @@
|
||||
This directory holds shared units for recipes, mostly shared configs
|
||||
|
||||
Recipes can reference shared configs via paths under configs/...
|
||||
@@ -0,0 +1,3 @@
|
||||
This directory holds model-agnostic general recipes
|
||||
|
||||
ptq/*: model-agnostic general Post Training Quantization recipes.
|
||||
@@ -0,0 +1,64 @@
|
||||
# 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.
|
||||
|
||||
metadata:
|
||||
recipe_type: ptq
|
||||
description: FP8 per-tensor weight and activation (W8A8), FP8 KV cache, max calibration.
|
||||
ptq_cfg:
|
||||
algorithm: max
|
||||
quant_cfg:
|
||||
'*input_quantizer':
|
||||
num_bits: e4m3
|
||||
axis:
|
||||
'*weight_quantizer':
|
||||
num_bits: e4m3
|
||||
axis:
|
||||
default:
|
||||
enable: false
|
||||
'*block_sparse_moe.gate*':
|
||||
enable: false
|
||||
'*linear_attn.conv1d*':
|
||||
enable: false
|
||||
'*lm_head*':
|
||||
enable: false
|
||||
'*mixer.conv1d*':
|
||||
enable: false
|
||||
'*mlp.gate.*':
|
||||
enable: false
|
||||
'*mlp.shared_expert_gate.*':
|
||||
enable: false
|
||||
'*output_layer*':
|
||||
enable: false
|
||||
'*proj_out.*':
|
||||
enable: false
|
||||
'*router*':
|
||||
enable: false
|
||||
output.*:
|
||||
enable: false
|
||||
nn.BatchNorm1d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.BatchNorm2d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.BatchNorm3d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.LeakyReLU:
|
||||
'*':
|
||||
enable: false
|
||||
'*[kv]_bmm_quantizer':
|
||||
num_bits: e4m3
|
||||
enable: true
|
||||
@@ -0,0 +1,72 @@
|
||||
# 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.
|
||||
|
||||
metadata:
|
||||
recipe_type: ptq
|
||||
description: NVFP4 MLP/MoE weight only (W4A16), FP8 KV cache, max calibration.
|
||||
ptq_cfg:
|
||||
algorithm: max
|
||||
quant_cfg:
|
||||
'*weight_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
'*input_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
default:
|
||||
enable: false
|
||||
'*block_sparse_moe.gate*':
|
||||
enable: false
|
||||
'*linear_attn.conv1d*':
|
||||
enable: false
|
||||
'*lm_head*':
|
||||
enable: false
|
||||
'*mixer.conv1d*':
|
||||
enable: false
|
||||
'*mlp.gate.*':
|
||||
enable: false
|
||||
'*mlp.shared_expert_gate.*':
|
||||
enable: false
|
||||
'*output_layer*':
|
||||
enable: false
|
||||
'*proj_out.*':
|
||||
enable: false
|
||||
'*router*':
|
||||
enable: false
|
||||
output.*:
|
||||
enable: false
|
||||
nn.BatchNorm1d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.BatchNorm2d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.BatchNorm3d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.LeakyReLU:
|
||||
'*':
|
||||
enable: false
|
||||
'*[kv]_bmm_quantizer':
|
||||
num_bits: e4m3
|
||||
enable: true
|
||||
@@ -0,0 +1,86 @@
|
||||
# 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.
|
||||
|
||||
metadata:
|
||||
recipe_type: ptq
|
||||
description: NVFP4 static weight and dynamic activation for all linear layers (W4A4), FP8 KV cache, max calibration.
|
||||
ptq_cfg:
|
||||
algorithm: max
|
||||
quant_cfg:
|
||||
'*mlp*weight_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
'*mlp*input_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
'*block_sparse_moe*weight_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
'*block_sparse_moe*input_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
default:
|
||||
enable: false
|
||||
'*block_sparse_moe.gate*':
|
||||
enable: false
|
||||
'*linear_attn.conv1d*':
|
||||
enable: false
|
||||
'*lm_head*':
|
||||
enable: false
|
||||
'*mixer.conv1d*':
|
||||
enable: false
|
||||
'*mlp.gate.*':
|
||||
enable: false
|
||||
'*mlp.shared_expert_gate.*':
|
||||
enable: false
|
||||
'*output_layer*':
|
||||
enable: false
|
||||
'*proj_out.*':
|
||||
enable: false
|
||||
'*router*':
|
||||
enable: false
|
||||
output.*:
|
||||
enable: false
|
||||
nn.BatchNorm1d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.BatchNorm2d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.BatchNorm3d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.LeakyReLU:
|
||||
'*':
|
||||
enable: false
|
||||
'*[kv]_bmm_quantizer':
|
||||
num_bits: e4m3
|
||||
enable: true
|
||||
@@ -0,0 +1,100 @@
|
||||
# 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.
|
||||
|
||||
metadata:
|
||||
recipe_type: ptq
|
||||
description: NVFP4 static weight and dynamic activation for all linear layers including output projections, FP8 KV cache, max calibration.
|
||||
ptq_cfg:
|
||||
algorithm: max
|
||||
quant_cfg:
|
||||
'*mlp*weight_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
'*mlp*input_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
'*block_sparse_moe*weight_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
'*block_sparse_moe*input_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
'*o_proj*weight_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
'*o_proj*input_quantizer':
|
||||
block_sizes:
|
||||
-1: 16
|
||||
type: dynamic
|
||||
scale_bits: e4m3
|
||||
num_bits: e2m1
|
||||
enable: true
|
||||
default:
|
||||
enable: false
|
||||
'*block_sparse_moe.gate*':
|
||||
enable: false
|
||||
'*linear_attn.conv1d*':
|
||||
enable: false
|
||||
'*lm_head*':
|
||||
enable: false
|
||||
'*mixer.conv1d*':
|
||||
enable: false
|
||||
'*mlp.gate.*':
|
||||
enable: false
|
||||
'*mlp.shared_expert_gate.*':
|
||||
enable: false
|
||||
'*output_layer*':
|
||||
enable: false
|
||||
'*proj_out.*':
|
||||
enable: false
|
||||
'*router*':
|
||||
enable: false
|
||||
output.*:
|
||||
enable: false
|
||||
nn.BatchNorm1d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.BatchNorm2d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.BatchNorm3d:
|
||||
'*':
|
||||
enable: false
|
||||
nn.LeakyReLU:
|
||||
'*':
|
||||
enable: false
|
||||
'*[kv]_bmm_quantizer':
|
||||
num_bits: e4m3
|
||||
enable: true
|
||||
@@ -47,6 +47,8 @@ dependencies = [
|
||||
"rich",
|
||||
"safetensors",
|
||||
"scipy",
|
||||
"PyYAML>=6.0",
|
||||
"omegaconf>=2.3.0"
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
@@ -122,6 +124,7 @@ include = ["modelopt*"]
|
||||
|
||||
[tool.setuptools.package-data]
|
||||
modelopt = ["**/*.h", "**/*.cpp", "**/*.cu"]
|
||||
modelopt_recipes = ["**/*.yml", "**/*.yaml"]
|
||||
|
||||
[tool.uv]
|
||||
managed = true
|
||||
|
||||
@@ -48,9 +48,12 @@ FP4_SVDQUANT_CFG["algorithm"] = {"method": "svdquant", "lowrank": 8}
|
||||
def get_awq_config(algorithm="awq_lite", block_size=8):
|
||||
config = copy.deepcopy(mtq.INT4_AWQ_CFG)
|
||||
config["quant_cfg"]["*weight_quantizer"]["block_sizes"] = {-1: block_size}
|
||||
if "algorithm" not in config or not isinstance(config["algorithm"], dict):
|
||||
config["algorithm"] = {}
|
||||
|
||||
config["algorithm"]["method"] = algorithm
|
||||
config["algorithm"]["debug"] = True
|
||||
if algorithm == "awq_clip":
|
||||
if algorithm == "awq_clip" and "alpha_step" in config["algorithm"]:
|
||||
config["algorithm"].pop("alpha_step")
|
||||
return config
|
||||
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
# 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.
|
||||
|
||||
@@ -0,0 +1,212 @@
|
||||
# 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.
|
||||
|
||||
"""Unit tests for modelopt.recipe.loader and modelopt.recipe.loader.load_config."""
|
||||
|
||||
import pytest
|
||||
|
||||
from modelopt.recipe.config import ModelOptPTQRecipe, RecipeType
|
||||
from modelopt.recipe.loader import load_config, load_recipe
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Static YAML fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
CFG_AB = """\
|
||||
a: 1
|
||||
b: 2
|
||||
"""
|
||||
|
||||
CFG_KEY_VAL = """\
|
||||
key: val
|
||||
"""
|
||||
|
||||
CFG_RECIPE_MISSING_TYPE = """\
|
||||
metadata:
|
||||
description: Missing recipe_type.
|
||||
ptq_cfg: {}
|
||||
"""
|
||||
|
||||
CFG_RECIPE_MISSING_PTQ_CFG = """\
|
||||
metadata:
|
||||
recipe_type: ptq
|
||||
"""
|
||||
|
||||
CFG_RECIPE_UNSUPPORTED_TYPE = """\
|
||||
metadata:
|
||||
recipe_type: unknown_type
|
||||
"""
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Directory-format YAML fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# load_config — basic behaviour
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_load_config_plain(tmp_path):
|
||||
"""A plain config is returned as-is."""
|
||||
(tmp_path / "cfg.yml").write_text(CFG_AB)
|
||||
assert load_config(tmp_path / "cfg.yml") == {"a": 1, "b": 2}
|
||||
|
||||
|
||||
def test_load_config_suffix_probe(tmp_path):
|
||||
"""load_config finds a .yml file when suffix is omitted from a string path."""
|
||||
(tmp_path / "mycfg.yml").write_text(CFG_KEY_VAL)
|
||||
assert load_config(str(tmp_path / "mycfg")) == {"key": "val"}
|
||||
|
||||
|
||||
def test_load_config_missing_file_raises(tmp_path):
|
||||
"""load_config raises ValueError for a path that does not exist."""
|
||||
with pytest.raises(ValueError, match="Cannot find config file"):
|
||||
load_config(str(tmp_path / "nonexistent"))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# load_recipe — built-in PTQ recipes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_load_recipe_builtin_with_suffix():
|
||||
"""load_recipe loads a built-in PTQ recipe given the full YAML path."""
|
||||
recipe = load_recipe("general/ptq/fp8_default-fp8_kv.yml")
|
||||
assert recipe.recipe_type == RecipeType.PTQ
|
||||
assert isinstance(recipe, ModelOptPTQRecipe)
|
||||
assert recipe.ptq_cfg
|
||||
|
||||
|
||||
def test_load_recipe_builtin_without_suffix():
|
||||
"""load_recipe resolves the .yml suffix automatically."""
|
||||
recipe = load_recipe("general/ptq/fp8_default-fp8_kv")
|
||||
assert recipe.recipe_type == RecipeType.PTQ
|
||||
|
||||
|
||||
def test_load_recipe_builtin_description():
|
||||
"""The description field is loaded from the YAML metadata."""
|
||||
recipe = load_recipe("general/ptq/fp8_default-fp8_kv.yml")
|
||||
assert isinstance(recipe.description, str)
|
||||
assert len(recipe.description) > 0
|
||||
|
||||
|
||||
_BUILTIN_PTQ_RECIPES = [
|
||||
"general/ptq/fp8_default-fp8_kv",
|
||||
"general/ptq/nvfp4_default-fp8_kv",
|
||||
"general/ptq/nvfp4_mlp_only-fp8_kv",
|
||||
"general/ptq/nvfp4_omlp_only-fp8_kv",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("recipe_path", _BUILTIN_PTQ_RECIPES)
|
||||
def test_load_recipe_all_builtins(recipe_path):
|
||||
"""Smoke-test: every built-in PTQ recipe loads without error and has ptq_cfg."""
|
||||
recipe = load_recipe(recipe_path)
|
||||
assert recipe.recipe_type == RecipeType.PTQ
|
||||
assert isinstance(recipe, ModelOptPTQRecipe)
|
||||
assert recipe.ptq_cfg
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# load_recipe — error cases
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_load_recipe_missing_raises(tmp_path):
|
||||
"""load_recipe raises ValueError for a path that doesn't exist."""
|
||||
with pytest.raises(ValueError):
|
||||
load_recipe(str(tmp_path / "does_not_exist.yml"))
|
||||
|
||||
|
||||
def test_load_recipe_missing_recipe_type_raises(tmp_path):
|
||||
"""load_recipe raises ValueError when metadata.recipe_type is absent."""
|
||||
bad = tmp_path / "bad.yml"
|
||||
bad.write_text(CFG_RECIPE_MISSING_TYPE)
|
||||
with pytest.raises(ValueError, match="recipe_type"):
|
||||
load_recipe(bad)
|
||||
|
||||
|
||||
def test_load_recipe_missing_ptq_cfg_raises(tmp_path):
|
||||
"""load_recipe raises ValueError when ptq_cfg is absent for a PTQ recipe."""
|
||||
bad = tmp_path / "bad.yml"
|
||||
bad.write_text(CFG_RECIPE_MISSING_PTQ_CFG)
|
||||
with pytest.raises(ValueError, match="ptq_cfg"):
|
||||
load_recipe(bad)
|
||||
|
||||
|
||||
def test_load_recipe_unsupported_type_raises(tmp_path):
|
||||
"""load_recipe raises ValueError for an unknown recipe_type."""
|
||||
bad = tmp_path / "bad.yml"
|
||||
bad.write_text(CFG_RECIPE_UNSUPPORTED_TYPE)
|
||||
with pytest.raises(ValueError, match="Unsupported recipe type"):
|
||||
load_recipe(bad)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# load_recipe — directory format
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_load_recipe_dir(tmp_path):
|
||||
"""load_recipe loads a recipe from a directory with recipe.yml + ptq_cfg.yml."""
|
||||
(tmp_path / "recipe.yml").write_text(
|
||||
"metadata:\n recipe_type: ptq\n description: Dir test.\n"
|
||||
)
|
||||
(tmp_path / "ptq_cfg.yml").write_text("algorithm: max\nquant_cfg: {}\n")
|
||||
recipe = load_recipe(tmp_path)
|
||||
assert recipe.recipe_type == RecipeType.PTQ
|
||||
assert recipe.description == "Dir test."
|
||||
assert recipe.ptq_cfg == {"algorithm": "max", "quant_cfg": {}}
|
||||
|
||||
|
||||
def test_load_recipe_dir_missing_recipe_raises(tmp_path):
|
||||
"""load_recipe raises ValueError when recipe.yml is absent from the directory."""
|
||||
(tmp_path / "ptq_cfg.yml").write_text("algorithm: max\nquant_cfg: {}\n")
|
||||
with pytest.raises(ValueError, match="recipe descriptor"):
|
||||
load_recipe(tmp_path)
|
||||
|
||||
|
||||
def test_load_recipe_dir_missing_ptq_cfg_raises(tmp_path):
|
||||
"""load_recipe raises ValueError when ptq_cfg.yml is absent from the directory."""
|
||||
(tmp_path / "recipe.yml").write_text("metadata:\n recipe_type: ptq\n")
|
||||
with pytest.raises(ValueError, match="ptq_cfg"):
|
||||
load_recipe(tmp_path)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# YAML recipe consistency — built-in general/ptq files match config.py dicts
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("yaml_path", "model_cfg_name", "kv_cfg_name"),
|
||||
[
|
||||
("general/ptq/fp8_default-fp8_kv.yml", "FP8_DEFAULT_CFG", "FP8_KV_CFG"),
|
||||
("general/ptq/nvfp4_default-fp8_kv.yml", "NVFP4_DEFAULT_CFG", "FP8_KV_CFG"),
|
||||
("general/ptq/nvfp4_mlp_only-fp8_kv.yml", "NVFP4_MLP_ONLY_CFG", "FP8_KV_CFG"),
|
||||
("general/ptq/nvfp4_omlp_only-fp8_kv.yml", "NVFP4_OMLP_ONLY_CFG", "FP8_KV_CFG"),
|
||||
],
|
||||
)
|
||||
def test_general_ptq_yaml_matches_config_dicts(yaml_path, model_cfg_name, kv_cfg_name):
|
||||
"""Each general/ptq YAML's merged quant_cfg matches the corresponding config.py dicts."""
|
||||
import modelopt.torch.quantization.config as qcfg
|
||||
|
||||
model_cfg = getattr(qcfg, model_cfg_name)
|
||||
kv_cfg = getattr(qcfg, kv_cfg_name)
|
||||
yaml_data = load_config(yaml_path)
|
||||
|
||||
ptq = yaml_data["ptq_cfg"]
|
||||
assert {**model_cfg["quant_cfg"], **kv_cfg["quant_cfg"]} == ptq["quant_cfg"]
|
||||
assert model_cfg["algorithm"] == ptq["algorithm"]
|
||||
Reference in New Issue
Block a user