Files
Wei-Ming Chen 1d392999b4 [OMNIML-5570] 2/2 Compose GEMM and KV-cache AutoQuant workflows (#2273)
### What does this PR do?

Type of change: new feature.

Follow-up to merged #2272. Adds composition of existing GEMM
quantization with KV-cache AutoQuantize:

- fixed FP8 GEMM PTQ followed by mixed-KV AutoQuantize;
- gradient-based NVFP4/FP8 GEMM AutoQuantize followed by independent
mixed-KV AutoQuantize;
- an optional `kv_auto_quantize` recipe stage with independent method,
constraints, candidates, and checkpoint path;
- ordered `hf_ptq.py` orchestration that keeps selected
weight/activation QDQ active while its calibration state remains frozen
during KV candidate calibration;
- fail-closed validation when a preceding stage leaves actual K/V
quantizers enabled; and
- unified export of a uniform-weight or mixed-weight checkpoint together
with the selected per-layer KV map.

The KV search still uses the public `mtq.auto_quantize(...,
constraints={"cost_model": "kv_cache", ...})` API from #2272. On a
converted model, the API preserves existing non-KV quantizers and
requires K/V to be disabled before search. Fresh-model behavior is
unchanged and starts from a deny-all quantizer baseline.

#### Why a follow-up field instead of a generic stage list?

This PR deliberately supports the two composition forms required by
`hf_ptq.py` without replacing the stable recipe schema. Existing recipes
already express a fixed `quantize` baseline plus one primary
`auto_quantize` search. A generic ordered `stages` list would require a
broader recipe/API migration, indexed checkpoint semantics, and
compatibility rules for arbitrary stage sequences. There is not yet a
demonstrated third search stage that justifies that surface-area change.

The two searches are not combined inside `mtq.auto_quantize`: each
invocation owns one search domain, constraint model, scoring method, and
resumable checkpoint. Their ordering and independent checkpoint paths
are orchestration concerns, while candidate calibration, scoring,
selection, and state application remain in the shared public API. A
general stage pipeline can be considered separately if more than this
one optional KV follow-up is needed.

Both solvers and scoring protocols are unchanged. The KV checkpoint
compatibility signature additionally fingerprints the preceding
quantizer configuration and calibrated state. Unsupported uniform-weight
plus mixed-KV exports record `kv_cache_deployment_supported: false` in
both ModelOpt and converted HF metadata.

### Usage

Fixed FP8 GEMM PTQ followed by KV AutoQuantize:

```bash
python examples/hf_ptq/hf_ptq.py \
  --pyt_ckpt_path Qwen/Qwen3-8B \
  --recipe general/auto_quantize/fp8_ptq_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits \
  --kv_auto_quantize_checkpoint /path/to/kv_autoquant.pth \
  --export_path /path/to/qwen3-8b-fp8-and-mixed-kv
```

Weight AutoQuantize followed by KV AutoQuantize:

```bash
python examples/hf_ptq/hf_ptq.py \
  --pyt_ckpt_path Qwen/Qwen3-8B \
  --recipe general/auto_quantize/nvfp4_fp8_gradient_then_kv_fp8_nvfp4_cast_kl_div_at_5p4bits \
  --auto_quantize_checkpoint /path/to/weight_autoquant.pth \
  --kv_auto_quantize_checkpoint /path/to/kv_autoquant.pth \
  --export_path /path/to/qwen3-8b-autoquant-and-mixed-kv
```

KV checkpoint resume requires identical preceding non-K/V quantizer
configuration and calibrated state. If rerunning the preceding stage
changes that state, use a new KV checkpoint path to recompute
sensitivities; configuration identity alone is insufficient to reuse the
scores safely.

### Testing

- Latest changed-area validation: 126 tests passed across `hf_ptq.py`
orchestration, KV checkpoint compatibility, export metadata, and HF
configuration conversion.
- A broader local run had 604 passes, one skip, and six failures: two
socket-binding failures under the sandbox and four local Transformers
API incompatibilities. This is not a full-suite pass.
- The fixed-PTQ→KV recipe executes end to end on a tiny offline Qwen
fixture.
- Public API coverage verifies that composed KV search preserves
preceding weight quantization and rejects enabled K/V state.
- Changed-file pre-commit hooks passed; the isolated recipe validator
also passed after dependency bootstrap.

### 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?: ✅
- Did you update Changelog?: ✅ (0.48.0 composition feature and KV
checkpoint flag deprecation)
- Did you get Claude approval on this PR?: ❌

### Additional Information

- This follow-up targets `main`, which contains merged #2272.
- `--auto_quantize_checkpoint` and `--kv_auto_quantize_checkpoint` are
intentionally separate because KV sensitivities depend on the preceding
GEMM state.
- Uniform-weight plus mixed-KV exports are for artifact inspection until
the runtime's uniform-weight ModelOpt configuration consumes
`kv_cache_quantized_layers`. Export emits an actionable warning and
records `kv_cache_deployment_supported: false`; this marker does not
itself add runtime support.


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

- **New Features**
- Added staged post-training quantization workflows for weights and KV
caches, including dedicated KV-cache checkpoints.
- Added FP8/NVFP4 recipes with configurable bit constraints and scoring.
  - KV-cache quantization now supports pre-quantized models.

- **Bug Fixes**
- Mixed weight and KV-cache quantization now exports with a warning
instead of failing.
- Improved validation and checkpoint compatibility for staged
configurations.
- Added safeguards for configurations without enabled weight quantizers.

- **Documentation**
- Clarified staged KV-cache workflows, checkpoint options, configuration
behavior, and unsupported deployment combinations.
- Documented deprecated legacy quantization options and their
replacement behavior.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: weimingc <17592131+meenchen@users.noreply.github.com>
2026-09-29 00:15:40 +00:00

517 lines
21 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""ModelOpt's pydantic BaseModel for recipes."""
from __future__ import annotations
import warnings
from enum import Enum
from typing import ClassVar, Literal
from pydantic import Field, field_validator, model_validator
from modelopt.torch.opt.config import ModeloptBaseConfig, ModeloptField
from modelopt.torch.quantization.config import QuantizeConfig # noqa: TC001
from modelopt.torch.speculative.config import DFlashConfig, EagleConfig, MedusaConfig
from modelopt.torch.speculative.plugins.hf_training_args import DataArguments as SpecDataArgs
from modelopt.torch.speculative.plugins.hf_training_args import ModelArguments as SpecModelArgs
from modelopt.torch.speculative.plugins.hf_training_args import (
TrainingArguments as SpecTrainingArgs,
)
__all__ = [
"RECIPE_TYPE_TO_CLASS",
"AutoQuantizeConfig",
"AutoQuantizeConstraints",
"AutoQuantizeCost",
"AutoQuantizeModuleSearchSpace",
"ModelOptAutoQuantizeRecipe",
"ModelOptDFlashRecipe",
"ModelOptEagleRecipe",
"ModelOptMedusaRecipe",
"ModelOptPTQRecipe",
"ModelOptRecipeBase",
"ModelOptSpeculativeRecipeBase",
"RecipeMetadataConfig",
"RecipeType",
]
class RecipeType(str, Enum):
"""List of recipe types. See ``RECIPE_TYPE_TO_CLASS`` at the bottom for the schema mapping."""
PTQ = "ptq"
AUTO_QUANTIZE = "auto_quantize"
SPECULATIVE_EAGLE = "speculative_eagle"
SPECULATIVE_DFLASH = "speculative_dflash"
SPECULATIVE_MEDUSA = "speculative_medusa"
# QAT = "qat" # Not implemented yet, will be added in the future.
_DEFAULT_RECIPE_DESCRIPTION = "Model optimization recipe."
class RecipeMetadataConfig(ModeloptBaseConfig):
"""YAML shape of the recipe metadata section."""
recipe_type: RecipeType | None = ModeloptField(
default=None,
title="Recipe type",
description="The type of the recipe (e.g. PTQ). **Deprecated** in recipe YAML: "
"the ``# modelopt-schema:`` comment naming the recipe's schema class already says "
"which kind it is -- and it is the same declaration that makes the file "
"``$import``-able -- so the class fills this in. Still read and still honoured, so "
"no existing recipe needs changing, but new recipes should leave it out -- including "
"in a directory-format recipe's ``metadata.yml``, which supports the same comment. "
"When both are present they must agree.",
)
description: str = ModeloptField(
default=_DEFAULT_RECIPE_DESCRIPTION,
title="Description",
description="Human-readable description of the recipe.",
)
def _metadata_field():
"""Build a metadata Pydantic field that defaults to the owning class's recipe type."""
return ModeloptField(
default={"description": _DEFAULT_RECIPE_DESCRIPTION},
title="Metadata",
description="Recipe metadata containing the recipe type and description.",
validate_default=True,
)
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}``.
"""
#: The kind of recipe this class *is*. Set on every concrete subclass; it is the
#: single source of truth for ``metadata.recipe_type``, which the validator below
#: fills in so a recipe file never has to repeat what its schema already states.
RECIPE_TYPE: ClassVar[RecipeType | None] = None
metadata: RecipeMetadataConfig = Field(
title="Metadata",
description="Recipe metadata containing the recipe type and description. "
"Required: a recipe without a ``metadata`` section is rejected so that a "
"recipe always says what it is for.",
)
@model_validator(mode="after")
def _resolve_recipe_type(self):
"""Fill ``metadata.recipe_type`` from the schema class, or reject a mismatch.
The schema class already determines the kind, so a recipe file that declares its
schema needs no ``recipe_type``. One that states it anyway must state the truth --
a silent disagreement between the two would make the file mean different things
to the loader and to a reader.
"""
if self.RECIPE_TYPE is None:
return self
if self.metadata.recipe_type is None:
self.metadata.recipe_type = self.RECIPE_TYPE
elif self.metadata.recipe_type != self.RECIPE_TYPE:
raise ValueError(
f"metadata.recipe_type is {self.metadata.recipe_type.value!r} but this recipe "
f"is a {type(self).__name__}, which is {self.RECIPE_TYPE.value!r}. Drop the "
"recipe_type (the schema declares it) or correct it."
)
return self
@property
def recipe_type(self) -> RecipeType:
"""Return the recipe type from metadata."""
assert self.metadata.recipe_type is not None, "recipe_type was not resolved"
return self.metadata.recipe_type
@property
def description(self) -> str:
"""Return the recipe description from metadata."""
return self.metadata.description
class ModelOptPTQRecipe(ModelOptRecipeBase):
"""Our config class for PTQ recipes."""
RECIPE_TYPE: ClassVar[RecipeType] = RecipeType.PTQ
quantize: QuantizeConfig = Field(
title="PTQ config",
description="PTQ config containing quant_cfg and algorithm. Required: a PTQ "
"recipe without a ``quantize`` section is rejected so that a missing section "
"can't silently fall back to the default INT8 config.",
)
# Named alias so a shared layer-pattern unit (e.g. configs/auto_quantize/units/base_disabled_layers)
# can declare ``modelopt-schema: modelopt.recipe.config.LayerPatternList`` and be spliced into a
# ``list[str]`` field — mirrors how base_disable_all is imported into a PTQ quant_cfg list.
LayerPatternList = list[str]
class AutoQuantizeCost(ModeloptBaseConfig):
"""Cost-model parameters (the ``cost`` sub-dict of ``mtq.auto_quantize`` constraints)."""
active_moe_expert_ratio: float | None = ModeloptField(
default=None,
title="Active MoE expert ratio",
description="Routed experts active per token, in (0, 1]. Used by the 'active_moe' cost model.",
)
@field_validator("active_moe_expert_ratio")
@classmethod
def _validate_active_moe_expert_ratio(cls, v: float | None) -> float | None:
if v is not None and not (0 < v <= 1):
raise ValueError(f"active_moe_expert_ratio must be in (0, 1], got {v}")
return v
class AutoQuantizeConstraints(ModeloptBaseConfig):
"""LP search constraints + cost model; matches the ``mtq.auto_quantize`` constraints dict."""
effective_bits: float = ModeloptField(
default=4.8,
title="Effective bits",
description=(
"Average storage-bits target for the selected cost model, in (0, 16]. Defaults to 4.8."
),
)
cost_model: Literal["weight", "active_moe", "kv_cache"] = ModeloptField(
default="weight",
title="Cost model",
description=(
"'weight' counts all weights equally; 'active_moe' scales routed-expert weights; "
"'kv_cache' accounts for paired K/V-cache storage."
),
)
cost: AutoQuantizeCost | None = ModeloptField(
default=None,
title="Cost-model parameters",
description="Extra cost-model parameters; omit for the 'weight' cost model.",
)
@field_validator("effective_bits")
@classmethod
def _validate_effective_bits(cls, v: float) -> float:
if not (0 < v <= 16):
raise ValueError(f"effective_bits must be in (0, 16], got {v}")
return v
@model_validator(mode="after")
def _validate_cost_settings(self):
if self.cost_model == "kv_cache" and self.cost is not None:
raise ValueError("KV-cache AutoQuant does not accept weight cost settings.")
return self
class AutoQuantizeModuleSearchSpace(ModeloptBaseConfig):
"""Candidate formats selectable for modules matching one or more name patterns."""
module_name_patterns: LayerPatternList = ModeloptField(
default=[],
title="Module name patterns",
description="Glob patterns matched against quantizable module names. A grouped AutoQuantize "
"decision must match a rule for every module in the group or for none of them.",
validate_default=True,
)
candidate_formats: list[QuantizeConfig] = ModeloptField(
default=[],
title="Module candidate quantization formats",
description="Formats selectable for matching modules. These override the top-level "
"candidate_formats for the matching AutoQuantize decision group.",
validate_default=True,
)
allow_no_quant: bool = ModeloptField(
default=True,
title="Allow no-quant selection",
description="Whether BF16/no-quant is selectable for matching modules. AutoQuantize keeps "
"an internal no-quant baseline for sensitivity scoring and cost normalization even when "
"this is false.",
)
@field_validator("module_name_patterns")
@classmethod
def _at_least_one_module_pattern(cls, v: list[str]) -> list[str]:
if not v:
raise ValueError("module_search_spaces requires at least 1 module_name_pattern")
return v
@field_validator("candidate_formats")
@classmethod
def _at_least_one_module_candidate(cls, v: list[QuantizeConfig]) -> list[QuantizeConfig]:
if not v:
raise ValueError("module_search_spaces requires at least 1 candidate_format")
return v
class AutoQuantizeConfig(ModeloptBaseConfig):
"""Schema for the ``auto_quantize`` block of an AutoQuantize recipe."""
constraints: AutoQuantizeConstraints = Field(
title="Search constraints + cost model",
description="LP budget and cost model.",
)
candidate_formats: list[QuantizeConfig] = ModeloptField(
default=[],
title="Candidate quantization formats",
description="Fallback per-layer search space for modules not matched by "
"module_search_spaces. Each entry is a full QuantizeConfig. BF16/no-quant is always an "
"implicit additional choice. Omit this field when the parent recipe supplies a fixed "
"quantize baseline and explicitly lists every searched family in module_search_spaces.",
validate_default=True,
)
module_search_spaces: list[AutoQuantizeModuleSearchSpace] = ModeloptField(
default=[],
title="Module-specific search spaces",
description="Optional per-module overrides for candidate formats and BF16/no-quant "
"selectability. Matching is performed after runtime-fusion grouping.",
)
auto_quantize_method: Literal["gradient", "kl_div"] = ModeloptField(
default="gradient",
title="Sensitivity scoring method",
description="'gradient' (Taylor + Fisher, needs labels) or 'kl_div' (no labels).",
)
score_size: int = ModeloptField(
default=128,
title="Scoring sample count",
description="Number of samples used for sensitivity scoring (divided by batch_size to get "
"the number of mtq scoring steps). Matches the former --auto_quantize_score_size.",
)
disabled_layers: LayerPatternList = ModeloptField(
default=[],
title="Search-excluded layer patterns",
description="Glob patterns; matching layers are excluded from the search (kept full precision).",
)
cost_excluded_layers: LayerPatternList = ModeloptField(
default=[],
title="Cost-excluded layer patterns",
description="Glob patterns excluded from the bit-budget accounting (cost_weight 0) — e.g. VL "
"vision towers. Distinct from disabled_layers: those are removed from the search; these still "
"get searched but don't count toward effective_bits. The two roles overlap but are independent.",
)
kv_cache: QuantizeConfig | None = ModeloptField(
default=None,
title="KV cache config (optional)",
description="QuantizeConfig applied as a uniform post-step; falls back to "
"the --kv_cache_qformat CLI flag when omitted.",
)
@model_validator(mode="after")
def _has_search_space(self):
if not self.candidate_formats and not self.module_search_spaces:
raise ValueError(
"auto_quantize requires candidate_formats or at least one module_search_spaces "
"entry. For uniform quantization, use a PTQ recipe instead."
)
if self.constraints.cost_model == "kv_cache":
if self.auto_quantize_method != "kl_div":
raise ValueError(
"KV-cache AutoQuant currently requires auto_quantize_method=kl_div."
)
if self.module_search_spaces:
raise ValueError(
"KV-cache AutoQuant uses one candidate space for all eligible attention "
"layers; module_search_spaces is not supported."
)
if self.kv_cache is not None:
raise ValueError(
"KV-cache AutoQuant candidate_formats replace the uniform kv_cache post-step."
)
if self.cost_excluded_layers:
raise ValueError(
"KV-cache AutoQuant does not support cost_excluded_layers; use "
"disabled_layers to exclude non-KV-cache modules from the search."
)
return self
class ModelOptAutoQuantizeRecipe(ModelOptRecipeBase):
"""Our config class for AutoQuantize recipes."""
RECIPE_TYPE: ClassVar[RecipeType] = RecipeType.AUTO_QUANTIZE
metadata: RecipeMetadataConfig = _metadata_field()
quantize: QuantizeConfig | None = ModeloptField(
default=None,
title="Fixed PTQ baseline",
description="Optional normal PTQ QuantizeConfig. A weight AutoQuantize stage uses it for "
"modules outside explicit module_search_spaces; a KV AutoQuantize stage applies it first "
"as the fixed GEMM weight/activation configuration.",
)
auto_quantize: AutoQuantizeConfig = Field(
title="AutoQuantize config",
description="AutoQuantize search configuration. Required.",
)
kv_auto_quantize: AutoQuantizeConfig | None = ModeloptField(
default=None,
title="Follow-up KV-cache AutoQuantize config",
description="Optional KV-cache search run after the primary weight AutoQuantize search.",
)
@model_validator(mode="after")
def _validate_fixed_and_searched_spaces(self):
primary_is_kv = self.auto_quantize.constraints.cost_model == "kv_cache"
if self.kv_auto_quantize is not None:
if primary_is_kv:
raise ValueError(
"kv_auto_quantize cannot follow an auto_quantize stage that already searches "
"the KV cache."
)
if self.kv_auto_quantize.constraints.cost_model != "kv_cache":
raise ValueError("kv_auto_quantize must use cost_model=kv_cache.")
if self.auto_quantize.kv_cache is not None:
raise ValueError(
"A weight AutoQuantize stage followed by kv_auto_quantize must omit the "
"uniform auto_quantize.kv_cache post-step."
)
has_fixed_baseline = self.quantize is not None
has_global_search = bool(self.auto_quantize.candidate_formats)
if not primary_is_kv and has_fixed_baseline and has_global_search:
raise ValueError(
"An AutoQuantize recipe with a fixed quantize baseline must omit top-level "
"auto_quantize.candidate_formats and explicitly list searched modules under "
"auto_quantize.module_search_spaces."
)
if not primary_is_kv and has_fixed_baseline and not self.auto_quantize.module_search_spaces:
raise ValueError(
"An AutoQuantize recipe with a fixed quantize baseline requires at least one "
"auto_quantize.module_search_spaces entry."
)
if not primary_is_kv and not has_fixed_baseline and not has_global_search:
raise ValueError(
"An AutoQuantize recipe without a fixed quantize baseline requires top-level "
"auto_quantize.candidate_formats for unmatched modules."
)
return self
class ModelOptSpeculativeRecipeBase(ModelOptRecipeBase):
"""Base class for speculative-decoding recipes.
Unlike PTQ, speculative-decoding is a training-time optimization: the draft head is trained
with HF Trainer. We therefore bundle ``model`` / ``data`` / ``training`` sections into the
recipe so a single YAML is the full experiment spec. Each section is a typed Pydantic model
(see :mod:`modelopt.torch.speculative.plugins.hf_training_args`) so field typos and bad
values are caught at recipe-load time; HF trainer fields pass through
``TrainingArguments`` via ``extra='allow'``.
"""
model: SpecModelArgs = ModeloptField(
default=SpecModelArgs(),
title="HF model args",
description="ModelArguments for the base HF model to train a draft head against.",
validate_default=True,
)
data: SpecDataArgs = ModeloptField(
default=SpecDataArgs(),
title="HF data args",
description="DataArguments for the training/offline dataset.",
validate_default=True,
)
training: SpecTrainingArgs = ModeloptField(
default=SpecTrainingArgs(),
title="HF training args",
description="Speculative-decoding extensions; HF trainer fields flow through as extras.",
validate_default=True,
)
class ModelOptEagleRecipe(ModelOptSpeculativeRecipeBase):
"""Our config class for EAGLE speculative decoding recipes."""
RECIPE_TYPE: ClassVar[RecipeType] = RecipeType.SPECULATIVE_EAGLE
metadata: RecipeMetadataConfig = _metadata_field()
eagle: EagleConfig = ModeloptField(
default=EagleConfig(),
title="EAGLE config",
description="EAGLE speculative decoding configuration.",
validate_default=True,
)
@model_validator(mode="after")
def _derive_eagle_offline(self) -> ModelOptEagleRecipe:
self.eagle.eagle_offline = self.data.mode != "online"
return self
@model_validator(mode="after")
def _warn_rope_vs_training_seq_len(self) -> ModelOptEagleRecipe:
orig_max_pos = self.eagle.eagle_export_rope_scaling.get("original_max_position_embeddings")
if orig_max_pos is not None and orig_max_pos != self.training.training_seq_len:
warnings.warn(
f"eagle.eagle_export_rope_scaling.original_max_position_embeddings ({orig_max_pos}) "
f"differs from training.training_seq_len ({self.training.training_seq_len}). "
f"This may affect long-context inference quality."
)
return self
class ModelOptDFlashRecipe(ModelOptSpeculativeRecipeBase):
"""Our config class for DFlash speculative decoding recipes."""
RECIPE_TYPE: ClassVar[RecipeType] = RecipeType.SPECULATIVE_DFLASH
metadata: RecipeMetadataConfig = _metadata_field()
dflash: DFlashConfig = ModeloptField(
default=DFlashConfig(),
title="DFlash config",
description="DFlash speculative decoding configuration.",
validate_default=True,
)
@model_validator(mode="after")
def _derive_dflash_offline(self) -> ModelOptDFlashRecipe:
# offline (dumped .pt) and streaming (hidden states via NIXL RDMA from a vLLM
# serve) both feed pre-computed base hidden states to the DFlash module, so
# both set dflash_offline. Only fully-online training runs the base model.
# Mirrors ModelOptEagleRecipe._derive_eagle_offline.
self.dflash.dflash_offline = self.data.mode != "online"
return self
class ModelOptMedusaRecipe(ModelOptSpeculativeRecipeBase):
"""Our config class for Medusa speculative decoding recipes."""
RECIPE_TYPE: ClassVar[RecipeType] = RecipeType.SPECULATIVE_MEDUSA
metadata: RecipeMetadataConfig = _metadata_field()
medusa: MedusaConfig = ModeloptField(
default=MedusaConfig(),
title="Medusa config",
description="Medusa speculative decoding configuration.",
validate_default=True,
)
# Single source of truth mapping YAML ``metadata.recipe_type`` to its schema class. The loader
# uses this for typed-list ``$import`` resolution; add a new entry when introducing a recipe.
RECIPE_TYPE_TO_CLASS: dict[RecipeType, type[ModelOptRecipeBase]] = {
RecipeType.PTQ: ModelOptPTQRecipe,
RecipeType.AUTO_QUANTIZE: ModelOptAutoQuantizeRecipe,
RecipeType.SPECULATIVE_EAGLE: ModelOptEagleRecipe,
RecipeType.SPECULATIVE_DFLASH: ModelOptDFlashRecipe,
RecipeType.SPECULATIVE_MEDUSA: ModelOptMedusaRecipe,
}