Source code for tensorrt_edgellm.config

# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
Model configuration parser for quantized LLM checkpoints.

Architecture fields are loaded with HuggingFace ``AutoConfig`` via
:func:`checkpoint_utils.load_checkpoint_config_dicts` (root + promoted LLM
dict; VL nested promotion applies to the LLM dict).  Quantization is merged from
``hf_quant_config.json`` or
embedded ``quantization_config`` via helpers in this module.

All model-specific feature flags (``attention_bias``, ``head_dim``, etc.)
are read from the resulting config dict.  The ``has_qk_norm`` flag is
auto-detected from the presence of ``q_norm`` weight keys in the checkpoint
index; no model-type string comparisons are used.

Supported quantization formats
-------------------------------
    fp16   - plain bfloat16/float16 weights (no quantization)
    fp8    - FP8 E4M3 per-tensor static quantization
    nvfp4  - NVFP4 per-group quantization with FP8 group scales
    nvfp4_a16 - weight-only NVFP4 (W4A16 / ModelOpt ``W4A16_NVFP4``)
    int4_awq            - AWQ INT4 group quantization (column-packed int32 checkpoints)
    int4_awq_modelopt   - W4A16_AWQ pre-packed uint8 ``[out//2, in]`` checkpoints
    int4_gptq           - GPTQ INT4 group quantization
    int8_sq             - INT8 SmoothQuant W8A8 per-channel
    mixed_precision     - per-layer mixed quantization (from ``hf_quant_config``)

Hybrid model support
--------------------
When ``config.json`` contains a ``layers_block_type`` list (e.g.
``["attention", "attention", "mamba", ...]``), a ``MambaConfig`` is parsed
and ``ModelConfig.layer_types`` reflects the per-layer block type.  Mamba
parameters (``mamba_num_heads``, ``mamba_head_dim``, ``ssm_state_size``,
``conv_dim``, ``conv_kernel``) are read from the config.  When ``conv_dim``
is absent (e.g. NemotronH-4B-BF16), it is derived from the ``conv1d.weight``
shape in the checkpoint to break the circular dependency with ``n_groups``.
"""

import fnmatch
import json
import math
import os
from dataclasses import asdict, dataclass, field, replace
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Tuple

if TYPE_CHECKING:
    import torch

from .checkpoint.checkpoint_utils import load_checkpoint_config_dicts
from .dflash import DFlashVersion, resolve_dflash_contract

# ---------------------------------------------------------------------------
# Quantization type constants
# ---------------------------------------------------------------------------

QUANT_FP16 = "fp16"
QUANT_FP8 = "fp8"
QUANT_MXFP8 = "mxfp8"
QUANT_NVFP4 = "nvfp4"
# Weight-only NVFP4 (W4A16): ModelOpt ``W4A16_NVFP4`` — 4-bit float weights,
# FP16 activations. Distinct from QUANT_NVFP4 (W4A4); routed to the dense/MoE
# Marlin FP16xE2M1 kernels.
QUANT_NVFP4_A16 = "nvfp4_a16"
QUANT_INT4_AWQ = "int4_awq"
QUANT_INT4_AWQ_MODELOPT = "int4_awq_modelopt"
QUANT_INT4_GPTQ = "int4_gptq"
QUANT_INT8_SQ = "int8_sq"
# ``quant_algo`` value in hf_quant_config for MIXED_PRECISION only.
# After parsing, :attr:`QuantConfig.quant_type` is always a concrete type
# (dominant algo); per-layer differences live in ``layer_overrides``.
QUANT_MIXED = "mixed_precision"

# Default RoPE base frequency (used when config omits rope_theta)
_DEFAULT_ROPE_THETA = 10000.0

# Model families whose per-head QK RMSNorm runs AFTER RoPE (HunYuan V1);
# the default (Qwen3 convention) normalizes before rotation.
_QK_NORM_POST_ROPE_MODEL_TYPES = frozenset({"hunyuan_v1_dense"})

# Model families that store the per-head QK norms under the HunYuan
# ``query_layernorm`` / ``key_layernorm`` key names. Kept in sync with the
# ``_hunyuan_key_remap`` dispatch in model.py — QK-norm detection must not
# fire for families whose loader would not remap these keys onto q/k_norm.
_QUERY_LAYERNORM_KEY_MODEL_TYPES = frozenset({"hunyuan_v1_dense"})

# Layer-type labels
LAYER_ATTN = "attention"
LAYER_MAMBA = "mamba"
LAYER_MLP = "mlp"
LAYER_GDN = "gdn"  # GatedDeltaNet linear attention (Qwen3.5)
LAYER_MOE = "moe"

_VALID_ATTENTION_LAYER_TYPES = ("sliding_attention", "full_attention")

# NemotronH ``hybrid_override_pattern`` / ``mtp_hybrid_override_pattern`` chars.
_HYBRID_PATTERN_MAP = {
    "M": LAYER_MAMBA,
    "-": LAYER_MLP,
    "*": LAYER_ATTN,
    "E": LAYER_MOE,
}

_DIFFUSION_GEMMA_MODEL_TYPES = frozenset({
    "diffusion_gemma",
    "diffusion_gemma_text",
    "diffusiongemma",
})
_DEFAULT_DIFFUSION_MAX_DENOISING_STEPS = 48


def _is_diffusion_gemma_model_type(model_type: str) -> bool:
    return str(model_type).lower() in _DIFFUSION_GEMMA_MODEL_TYPES


def _is_diffusion_gemma_config(root_dict: Dict[str, Any],
                               llm_dict: Dict[str, Any]) -> bool:
    model_type = str(
        root_dict.get("model_type", llm_dict.get("model_type", ""))).lower()
    architectures_value = (root_dict.get("architectures")
                           or llm_dict.get("architectures") or [])
    if isinstance(architectures_value, str):
        architectures = [architectures_value]
    else:
        architectures = [str(x) for x in architectures_value]
    return (_is_diffusion_gemma_model_type(model_type)
            or any("DiffusionGemma" in arch for arch in architectures))


def _is_gemma4_model_type(model_type: str) -> bool:
    model_type = str(model_type)
    return (model_type.startswith("gemma4")
            or _is_diffusion_gemma_model_type(model_type))


def _is_gemma4_assistant_model_type(model_type: str) -> bool:
    return str(model_type) in ("gemma4_assistant", "gemma4_unified_assistant")


def _is_muse_glimmer_model_type(model_type: str) -> bool:
    """Return whether a model type belongs to Muse-Glimmer."""
    return str(model_type) in ("muse_glimmer", "muse_glimmer_text")


def _check_num_attention_heads(num_attn_heads: int) -> None:
    if num_attn_heads <= 0:
        raise ValueError("num_attention_heads must be a positive integer, "
                         f"got {num_attn_heads!r}")


def _get_rope_theta(llm_dict: Dict[str, Any]) -> float:
    """Extract rope_theta from config dict.

    Some models (e.g. Qwen3) store rope_theta inside ``rope_scaling`` or
    ``rope_parameters`` rather than as a top-level key, so fall back to those
    nested dicts before returning the default 10 000.
    """
    if llm_dict.get("rope_theta") is not None:
        return float(llm_dict["rope_theta"])
    for key in ("rope_scaling", "rope_parameters"):
        nested = llm_dict.get(key)
        if not isinstance(nested, dict):
            continue
        if nested.get("rope_theta") is not None:
            return float(nested["rope_theta"])
        for attention_type in ("full_attention", "sliding_attention"):
            attention_params = nested.get(attention_type)
            if (isinstance(attention_params, dict)
                    and attention_params.get("rope_theta") is not None):
                return float(attention_params["rope_theta"])
    return _DEFAULT_ROPE_THETA


def _normalize_rope_scaling_for_config(
        rope_params: Optional[Dict[str, Any]]) -> Optional[dict]:
    """Normalize one RoPE parameter block for runtime config export."""
    if not isinstance(rope_params, dict):
        return None
    normalized = dict(rope_params)
    rope_type = normalized.get("rope_type", normalized.get("type"))
    if normalized.get("mrope_section") is not None:
        normalized["rope_type"] = "mrope"
        normalized["type"] = "mrope"
    elif rope_type is not None:
        normalized.setdefault("rope_type", rope_type)
        normalized.setdefault("type", rope_type)
    return normalized


def _select_rope_scaling(llm_dict: Dict[str, Any]) -> Optional[dict]:
    """Select the single-RoPE fallback block from raw config metadata."""
    for key in ("rope_scaling", "rope_parameters"):
        nested = llm_dict.get(key)
        if not isinstance(nested, dict):
            continue
        if isinstance(nested.get("full_attention"), dict):
            return _normalize_rope_scaling_for_config(nested["full_attention"])
        if isinstance(nested.get("sliding_attention"), dict):
            return _normalize_rope_scaling_for_config(
                nested["sliding_attention"])
        return _normalize_rope_scaling_for_config(nested)
    return None


def _runtime_rope_config_from_params(llm_dict: Dict[str, Any],
                                     rope_params: Dict[str, Any]) -> dict:
    """Build one runtime RoPE config block from raw checkpoint metadata."""
    out = {
        "rope_theta":
        float(
            rope_params.get("rope_theta",
                            llm_dict.get("rope_theta", _DEFAULT_ROPE_THETA))),
        "rope_scaling":
        _normalize_rope_scaling_for_config(rope_params),
        "partial_rotary_factor":
        float(
            rope_params.get("partial_rotary_factor",
                            llm_dict.get("partial_rotary_factor", 1.0))),
        "max_position_embeddings":
        int(llm_dict.get("max_position_embeddings", 4096)),
    }
    return out


def _get_dual_rope_configs(llm_dict: Dict[str, Any]) -> dict[str, dict]:
    """Return explicit sliding/full runtime RoPE configs when present."""
    rope_parameters = llm_dict.get("rope_parameters")
    layer_types = {
        str(layer_type)
        for layer_type in llm_dict.get("layer_types", [])
    }
    if not isinstance(rope_parameters, dict):
        return {}
    if not {"sliding_attention", "full_attention"} <= layer_types:
        return {}

    sliding_params = rope_parameters.get("sliding_attention")
    full_params = rope_parameters.get("full_attention")
    if not isinstance(sliding_params, dict) or not isinstance(
            full_params, dict):
        return {}
    return {
        "sliding_rope_config":
        _runtime_rope_config_from_params(llm_dict, sliding_params),
        "full_rope_config":
        _runtime_rope_config_from_params(llm_dict, full_params),
    }


def _parse_attention_layer_types(config: dict, num_hidden_layers: int,
                                 model_type: str) -> List[str]:
    """Preserve per-layer sliding/full attention labels when required."""
    raw = config.get("layer_types")
    if not (_is_gemma4_model_type(model_type)
            or _is_muse_glimmer_model_type(model_type)):
        return []

    if not isinstance(raw, list):
        raise ValueError(
            f"{model_type} requires layer_types with one sliding/full "
            "attention entry per layer.")
    if len(raw) != num_hidden_layers:
        raise ValueError(
            f"{model_type} layer_types length must match num_hidden_layers: "
            f"{len(raw)} vs {num_hidden_layers}.")

    attention_layer_types: List[str] = []
    for layer_idx, layer_type in enumerate(raw):
        layer_type = str(layer_type)
        if layer_type not in _VALID_ATTENTION_LAYER_TYPES:
            raise ValueError(
                f"{model_type} layer_types entries must be sliding_attention "
                "or full_attention; "
                f"got {layer_type!r} at layer {layer_idx}.")
        attention_layer_types.append(layer_type)
    return attention_layer_types


def _get_attention_scaling(llm_dict: Dict[str, Any], head_dim: int,
                           default_val: float) -> float:
    """Resolve the absolute multiplier applied to QK^T before softmax.

    Args:
        llm_dict: Checkpoint configuration containing an optional supported
            attention-scale alias.
        head_dim: Per-head query/key dimension.
        default_val: Required model-family fallback.

    Returns:
        The finite, positive attention scale supplied by the checkpoint, or
        the caller-selected fallback.

    Raises:
        ValueError: If an explicit checkpoint scale is not finite and
            positive, or if ``head_dim`` is not positive.
    """
    if head_dim <= 0:
        raise ValueError(f"head_dim must be positive; got {head_dim}")

    for key in ("attention_scaling", "qk_scale", "scaling"):
        if llm_dict.get(key) is not None:
            attention_scale = float(llm_dict[key])
            if not math.isfinite(attention_scale) or attention_scale <= 0.0:
                raise ValueError(
                    f"{key} must be finite and positive; got {llm_dict[key]!r}"
                )
            return attention_scale

    attention_scale = float(default_val)
    if not math.isfinite(attention_scale) or attention_scale <= 0.0:
        raise ValueError(
            f"default_val must be finite and positive; got {default_val!r}")
    return attention_scale


def _get_rms_norm_eps(llm_dict: Dict[str, Any], model_type: str) -> float:
    if str(model_type).lower().startswith("nemotron_h"):
        return llm_dict.get(
            "rms_norm_eps",
            llm_dict.get("norm_eps", llm_dict.get("layer_norm_epsilon", 1e-6)))

    return llm_dict.get("rms_norm_eps", 1e-6)


def _get_embedding_scale(llm_dict: Dict[str, Any], model_type: str,
                         hidden_size: int) -> float:
    """Return the scale folded into runtime token embeddings."""
    for key in ("embedding_scale", "embed_scale", "scalar_embed_scale"):
        if llm_dict.get(key) is not None:
            return float(llm_dict[key])

    if _is_gemma4_model_type(model_type):
        return math.sqrt(float(hidden_size))

    return 1.0


def _get_has_value_norm(llm_dict: Dict[str, Any], model_type: str) -> bool:
    """Return whether attention values use Gemma-style RMSNorm."""
    for key in ("has_value_norm", "has_v_norm", "value_norm"):
        if llm_dict.get(key) is not None:
            return bool(llm_dict[key])

    return _is_gemma4_model_type(model_type)


@dataclass
class Mapping:
    """Tensor-parallel placement for one exported or loaded rank.

    ``tp_size`` / ``tp_rank`` are the only supported model-sharding
    coordinates in the current Edge-LLM multi-device flow. Future CP/EP/PP/DP
    work should add its own mapping fields together with matching export,
    runtime, and validation support.
    """
    world_size: int = 1
    rank: int = 0
    tp_size: int = 1
    tp_rank: int = 0


@dataclass
class ActionConfig:
    """Action expert hyper-parameters for Alpamayo models.

    Used only for the action expert ONNX export and its sidecar config.json.
    These fields do NOT feed into the LLM config.json.
    """

    rope_theta: float = 5_000_000.0
    mrope_section: List[int] = field(default_factory=lambda: [24, 20, 20])
    mrope_interleaved: bool = True
    num_hidden_layers: int = 0
    num_attention_heads: int = 0
    num_key_value_heads: int = 0
    head_dim: int = 128
    attention_scaling: Optional[float] = None
    hidden_size: int = 0
    intermediate_size: int = 0
    rms_norm_eps: float = 1e-6
    num_traj_tokens: int = 1000
    traj_token_start: int = 0
    n_diffusion_tokens: int = 64
    in_proj_hidden_size: int = 512
    in_proj_num_enc_layers: int = 2
    in_proj_max_freq: float = 100.0
    in_proj_num_fourier_feats: int = 20

    def __post_init__(self) -> None:
        """Resolve and validate attention scaling for direct construction.

        Raises:
            ValueError: If the head dimension or explicit scale is invalid.
        """
        scale_config = ({} if self.attention_scaling is None else {
            "attention_scaling": self.attention_scaling
        })
        self.attention_scaling = _get_attention_scaling(
            scale_config, self.head_dim, 1.0 / (float(self.head_dim)**0.5))


_DEFAULT_QUANTIZE_ACTIVATIONS = True


def set_default_quantize_activations(value: bool) -> None:
    """Set :attr:`QuantConfig.quantize_activations` for configs parsed later.

    The export CLI calls this once from its argument parsing. A whole export can
    build several QuantConfigs (backbone plus any draft model), all of them
    parsed from their checkpoints afterwards, so one assignment covers the run
    without threading the flag through every ``_export_*`` entry point.
    """
    global _DEFAULT_QUANTIZE_ACTIVATIONS
    _DEFAULT_QUANTIZE_ACTIVATIONS = value


[docs] @dataclass class QuantConfig: """Quantization parameters extracted from the checkpoint config.""" quant_type: str = QUANT_FP16 # group_size: 1 = per-tensor/per-channel, 16 for NVFP4, 128 for AWQ group_size: int = 1 # GPTQ checkpoints are not consistent about whether qzeros stores the # actual zero point or (zero point - 1). The loader uses: # actual_zero = stored_zero + gptq_zero_point_offset. gptq_zero_point_offset: int = 1 # kv_cache_quant: "fp8" when KV-cache is quantised, None otherwise kv_cache_quant: Optional[str] = None # visual_mha_quant: "fp8" when the ViT visual attention (Q*K^T / P*V # matmuls) is quantised, None otherwise. Orthogonal to kv_cache_quant # (which only gates the LLM KV cache). visual_mha_quant: Optional[str] = None # module names excluded from quantisation (typically ["lm_head"]) excluded: List[str] = field(default_factory=list) # Per-layer quant type overrides for MIXED_PRECISION checkpoints. # Maps a module name (e.g. "lm_head") to a quant-type string. # make_linear() uses module_name together with ``excluded`` and (for # lm_head) ``ModelConfig.tie_word_embeddings`` to pick FP16 vs overrides. layer_overrides: dict = field(default_factory=dict) # Per-layer group sizes use the same normalized names as layer_overrides. layer_group_sizes: dict = field(default_factory=dict) # True when quant_algo is MIXED_PRECISION: unlisted modules are FP16. is_mixed_precision: bool = False # False exports quantized dense Linears without the activation Q-DQ pair, # leaving ``MatMul(fp16_activation, DQ(quantized_weight))``. Not a choice of # kernel: the weights stay in the checkpoint's format and TensorRT is left # to dequantize them, so this trades a large amount of decode throughput for # the accuracy of an unquantized activation. The default reproduces the # checkpoint's own recipe (W4A4 / W8A8); see # :func:`set_default_quantize_activations`. quantize_activations: bool = field( default_factory=lambda: _DEFAULT_QUANTIZE_ACTIVATIONS) @property def is_quantized(self) -> bool: return self.quant_type != QUANT_FP16 @property def uses_nvfp4_weights(self) -> bool: """True if any linear uses NVFP4 weights (dominant quant or layer override).""" if self.quant_type == QUANT_NVFP4: return True return any(v == QUANT_NVFP4 for v in self.layer_overrides.values()) @property def uses_mxfp8_weights(self) -> bool: """True if any linear uses MXFP8 weights (dominant quant or layer override).""" if self.quant_type == QUANT_MXFP8: return True return any(v == QUANT_MXFP8 for v in self.layer_overrides.values())
def module_quant_type(module_name: str, model_config: "ModelConfig") -> str: """Return the effective quant type ``make_linear`` will pick for *module_name*. Single source of truth for "what precision does this module's Linear end up at?" — used by :func:`make_linear` to pick the Linear class, and by the ONNX exporter to validate that LM-head externalization is only requested for an fp16 head. Lookup matches the names used elsewhere: ``ModelConfig.quant.excluded`` and ``layer_overrides`` keys are already normalized (VL prefixes / ``model.`` stripped) when parsed, and callers pass the same short ``module_name`` that ``make_linear`` receives (e.g. ``"lm_head"``). """ quant = model_config.quant if module_name and any( fnmatch.fnmatchcase(module_name, pattern) for pattern in quant.excluded): return QUANT_FP16 # Tied lm_head with no explicit override and an unquantized backbone has # no separate lm_head.weight in the checkpoint; treat it as fp16 so the # weight can be cloned from embed_tokens after loading. if (module_name == "lm_head" and model_config.tie_word_embeddings and "lm_head" not in quant.layer_overrides and quant.quant_type == QUANT_FP16): return QUANT_FP16 quant_type = quant.quant_type if module_name and quant.layer_overrides: fallback = QUANT_FP16 if quant.is_mixed_precision else quant_type quant_type = quant.layer_overrides.get(module_name, fallback) return quant_type def module_quant_group_size(module_name: str, model_config: "ModelConfig") -> int: """Return the checkpoint group size for a quantized linear module.""" quant = model_config.quant return int(quant.layer_group_sizes.get(module_name, quant.group_size)) @dataclass class MambaConfig: """Mamba-layer hyper-parameters for hybrid models.""" num_heads: int # mamba_num_heads head_dim: int # mamba_head_dim ssm_state_size: int # ssm_state_size conv_dim: int # conv_dim (total convolution channel count) conv_kernel: int # conv1d kernel size (default 4) n_groups: int # number of SSM groups (derived if not in config) @property def intermediate_size(self) -> int: """num_heads x head_dim - the Mamba intermediate feature dimension.""" return self.num_heads * self.head_dim @dataclass class GdnConfig: """GatedDeltaNet (GDN) hyper-parameters for Qwen3.5 hybrid models. GDN layers use a gated delta-net linear attention mechanism with fused QKV projection through causal conv1d. """ num_key_heads: int # linear_num_key_heads num_value_heads: int # linear_num_value_heads key_head_dim: int # linear_key_head_dim value_head_dim: int # linear_value_head_dim conv_kernel: int # linear_conv_kernel_dim (default 4) @property def key_dim(self) -> int: """Total key dimension (num_key_heads * key_head_dim).""" return self.num_key_heads * self.key_head_dim @property def value_dim(self) -> int: """Total value dimension (num_value_heads * value_head_dim).""" return self.num_value_heads * self.value_head_dim @property def conv_dim(self) -> int: """Total conv1d channel count: key + key + value (QKV fused).""" return self.key_dim + self.key_dim + self.value_dim @dataclass class DiffusionConfig: """Block-diffusion generation parameters parsed at export time.""" diffusion_family: str = "uniform_renoise" canvas_length: int = 256 max_denoising_steps: int = _DEFAULT_DIFFUSION_MAX_DENOISING_STEPS t_max: float = 0.8 t_min: float = 0.4 sampler_type: str = "entropy_bound" entropy_bound: float = 0.1 entropy_threshold: float = 0.005 stability_window: int = 2 self_conditioning_enabled: bool = True self_conditioning_repr: str = "embeds" supported_modalities: List[str] = field(default_factory=lambda: ["text"]) @classmethod def from_hf( cls, hf_config: Dict[str, Any], gen_config: Optional[Dict[str, Any]] = None) -> "DiffusionConfig": gen_config = gen_config or {} sampler_cfg = gen_config.get("sampler_config", {}) or {} sampler_name = str(sampler_cfg.get("_cls_name", "EntropyBound")) sampler_type = ("entropy_bound" if "entropybound" in sampler_name.lower().replace( "_", "") else sampler_name.lower()) return cls( canvas_length=int( hf_config.get("canvas_length", gen_config.get("canvas_length", 256))), max_denoising_steps=int( gen_config.get("max_denoising_steps", _DEFAULT_DIFFUSION_MAX_DENOISING_STEPS)), t_max=float(gen_config.get("t_max", 0.8)), t_min=float(gen_config.get("t_min", 0.4)), sampler_type=sampler_type, entropy_bound=float( sampler_cfg.get("entropy_bound", gen_config.get("entropy_bound", 0.1))), entropy_threshold=float( gen_config.get("entropy_threshold", gen_config.get("confidence_threshold", 0.005))), stability_window=int( gen_config.get("stability_window", gen_config.get("stability_threshold", 2))), ) def to_dict(self) -> Dict[str, Any]: return asdict(self)
[docs] @dataclass class ModelConfig: """Flat model hyper-parameter config consumed by module builders.""" # ------------------------------------------------------------------ arch model_type: str # HF architecture name, e.g. "qwen3", "llama" hidden_size: int num_hidden_layers: int num_attention_heads: int num_key_value_heads: int intermediate_size: int head_dim: int rms_norm_eps: float vocab_size: int rope_theta: float max_position_embeddings: int default_attention_scale: float # RoPE scaling config (e.g. {"type": "dynamic", "factor": 2.0} for Qwen2). # None means no scaling (standard RoPE). rope_scaling: Optional[dict] = None # For longrope: original context window size before extension. original_max_position_embeddings: Optional[int] = None # Fraction of head_dim used for RoPE (e.g. 0.75 for phi3/phi4, 1.0 for most others). partial_rotary_factor: float = 1.0 # Gemma4: head_dim for global (full_attention) layers (0 = same as head_dim) global_head_dim: int = 0 # Gemma4 E4B: num_key_value_heads for global (full_attention) layers (0 = same as num_key_value_heads) num_global_key_value_heads: int = 0 # Hidden activation name used by architecture-specific auxiliary modules. hidden_activation: str = "silu" # CodePredictor: RVQ code groups (lm_heads count = num_code_groups - 1). num_code_groups: int = 0 # Optional explicit RoPE configs for mixed sliding/full attention stacks. sliding_rope_config: Optional[dict] = None full_rope_config: Optional[dict] = None # ------------------------------------------ model-family feature flags # Per-head RMSNorm after Q and K projections. # Auto-detected from checkpoint key names; not inferred from model_type. has_qk_norm: bool = False # QK-norm order relative to RoPE. False (Qwen3 convention): norm then # rotate. True (HunYuan V1): rotate then norm — gamma placement differs, # so the attention plugin must apply the norm after rotation. qk_norm_post_rope: bool = False # Per-head RMSNorm after V projection. Gemma4 stores this norm without # learned weights, so it is selected from config metadata instead of # checkpoint key names. has_value_norm: bool = False # Bias on q/k/v projections. Read from config.json "attention_bias". attention_bias: bool = False # Explicit multiplicative scale applied to QK^T before softmax. attention_scaling: Optional[float] = None # Gemma4 full/global attention can use a different per-head dimension # from sliding attention. global_head_dim: Optional[int] = None # Gemma4 K=V full/global attention can use a different KV head count from # sliding attention. num_global_key_value_heads: Optional[int] = None # Gemma4 full/global attention reuses k_proj(hidden_states) as the value # projection source when enabled. attention_k_eq_v: bool = False # Per-Q-head learned attention sink: extra logit merged into softmax denominator. attention_sink_bias: bool = False # DiffusionGemma uses one shared backbone with phase-dependent layer scalars. encoder_layer_scalars: List[float] = field(default_factory=list) decoder_layer_scalars: List[float] = field(default_factory=list) self_conditioning_size: int = 0 diffusion: Optional[DiffusionConfig] = None # Multiplicative scale applied by the HF embedding module. embedding_scale: float = 1.0 # Final logit softcapping: tanh(logits/cap)*cap. None = disabled. final_logit_softcapping: Optional[float] = None # Muse-Glimmer: pre-tanh logit multiplier (applied before the softcap). output_multiplier: float = 1.0 # Muse-Glimmer: eps for the post-attention / post-FFN sandwich norms # (None -> fall back to rms_norm_eps). post_norm_eps: Optional[float] = None # Weight dtype in the checkpoint torch_dtype: str = "bfloat16" # When True, embed_tokens and lm_head share the same weight tensor tie_word_embeddings: bool = False # Sliding window attention size; -1 means no sliding window. sliding_window_size: int = -1 # Skip-softmax (BLASST) calibrated scale factor S; 0.0 disables the feature. skip_softmax_scale_factor: float = 0.0 # Gemma4 Unified 12B+: image placeholder runs use block-causal # attention during prefill (bidirectional inside each contiguous vision # run, causal everywhere else). Audio placeholders remain causal. use_vision_bidirectional_attention: bool = False # ------------------------------------------ per-layer block types # One entry per hidden layer: LAYER_ATTN, LAYER_MAMBA, LAYER_MLP, or LAYER_MOE. layer_types: List[str] = field(default_factory=list) # Original attention labels for attention layers: full_attention or sliding_attention. attention_layer_types: List[str] = field(default_factory=list) # ------------------------------------------ multimodal deepstack (VL) # Number of deepstack visual embedding tensors injected into the first N # hidden layers. Prefer ``vision_config.deepstack_visual_indexes`` length on # the root config when present; else ``num_deepstack_features`` or fallback. num_deepstack_features: int = 0 # ----------------------------------------- Qwen3-Omni emitted-tensor layer # Read by the Qwen3-Omni-specific ``Transformer`` subclasses (dense and # MoE) to decide which tensor to expose via ``emitted_hidden_states`` # (consumed by their ``CausalLM`` wrappers when # ``emit_hidden_states = True``): # # * ``accept_hidden_layer >= 1`` (and ≤ num_hidden_layers): pre-norm # output of decoder layer ``accept_hidden_layer - 1``. Matches HF's # ``outputs.hidden_states[k]`` convention where ``k`` denotes "after # ``k`` decoder layers" (k=0 is inputs_embeds, k=N is the last layer # pre-norm). Used by Qwen3-Omni Thinker → Talker: Talker # consumes ``thinker.hidden_states[accept_hidden_layer]``. # # * Default ``-1``: post-final-norm output (= ``model.norm(last_layer)``). # Used by Qwen3-Omni Talker → CodePredictor: HF reads # ``hidden_states[0][-1]`` which resolves to the post-norm tensor. accept_hidden_layer: int = -1 # -------------------------------------------------- quantization config quant: QuantConfig = field(default_factory=QuantConfig) # ------------------------------------------ mamba / hybrid config mamba_cfg: Optional[MambaConfig] = None # ------------------------------------------ gdn / hybrid config gdn_cfg: Optional[GdnConfig] = None # ------------------------------------------ gated attention (Qwen3.5) attn_output_gate: bool = False # ------------------------------------------ MTP config mtp_num_hidden_layers: Optional[int] = None mtp_use_dedicated_embeddings: bool = False # Nemotron-H MTP: the draft module is a hybrid stack whose layer types come # from this pattern (e.g. "*E" -> [attention, MoE]); ``mtp_num_hidden_layers`` # is then its length. ``num_nextn_predict_layers`` (the count of MTP prediction # modules) is folded in during parsing. mtp_hybrid_override_pattern: Optional[str] = None # The same draft stack, resolved from either the pattern above or the # ``mtp_layers_block_type`` list. mtp_layer_types: List[str] = field(default_factory=list) # When True, the standard CausalLM is exported as the MTP base model variant # with tree-attention inputs (attention_mask, attention_pos_id) and # an extra hidden_states output. mtp_base: bool = False # ------------------------------------------ Gemma4 MTP config root_model_type: str = "" raw_layer_types: List[str] = field(default_factory=list) rope_parameters: Optional[dict] = None backbone_hidden_size: int = 0 gemma4_mtp_base: bool = False gemma4_mtp_draft: bool = False assistant_hidden_size: int = 0 shares_target_kv: bool = False has_own_kv_cache: bool = True constant_draft_positions: bool = False returns_feedback_hidden: bool = False use_ordered_embeddings: bool = False num_centroids: int = 0 centroid_intermediate_top_k: int = 0 sparse_logits_enabled: bool = False kv_sharing_map: List[dict] = field(default_factory=list) # When True, MTP base export also exposes DDTree parent/depth metadata # for Qwen3.5 hybrid causal-conv/GDN tree-state execution (MTP tree drafting). mtp_tree_base: bool = False # ------------------------------------------ EAGLE3 draft config draft_vocab_size: Optional[int] = None target_hidden_size: Optional[int] = None is_eagle3_draft_flag: bool = False eagle3_target_layer_ids: List[int] = field(default_factory=list) # ------------------------------------------ EAGLE3 base config # When True, the standard CausalLM is exported as an EAGLE3 base model # with tree-attention inputs (attention_mask, attention_pos_id) and # an extra hidden_states output (concatenated from 3 selected layers). eagle_base: bool = False # ------------------------------------------ DFlash config # When True, export the standard Qwen3.5 model as the DFlash base with # tree-attention verify inputs and multi-layer hidden_states output. dflash_base: bool = False # When True, DFlash base export also exposes DDTree parent/depth metadata # for Qwen3.5 hybrid causal-conv/GDN tree-state execution. dflash_tree_base: bool = False is_dflash_draft_flag: bool = False dflash_target_layer_ids: List[int] = field(default_factory=list) dflash_block_size: int = 16 dflash_mask_token_id: int = 248070 dflash_version: DFlashVersion = DFlashVersion.V1 # Run the fc feature projector at the checkpoint's native precision (e.g. # NVFP4) instead of the default dense-FP16 + FP32 projection. Enabled only # for targets measured to keep target-hidden well inside FP16 range # (Nemotron-3.5). Qwen3-8B keeps the FP32 guard (target-hidden ~abs 2e4). dflash_fc_native_precision: bool = False # DFlash2 is a distinct linear-path proposal architecture. Its checkpoint # block size is the runtime default; an engine may profile a larger block. dflash2_target_layer_ids: List[int] = field(default_factory=list) dflash2_block_size: int = 8 dflash2_mask_token_id: int = 248070 dflash2_is_causal: bool = False dflash2_conv_kernel_size: int = 2 dflash2_conv_group_size: int = 16 dflash2_selector_rank: int = 256 dflash2_selector_top_k: int = 16 # ------------------------------------------ JetSpec config # JetSpec uses the DFlash/DDTree cached-draft contract with causal proposal # attention inside the draft block. The DFlash-prefixed fields are still # populated for the shared DFlashDraftModel ONNX implementation. jetspec_base: bool = False jetspec_tree_base: bool = False is_jetspec_draft_flag: bool = False jetspec_target_layer_ids: List[int] = field(default_factory=list) jetspec_block_size: int = 16 jetspec_mask_token_id: int = 151669 jetspec_causal_head: bool = False # ------------------------------------------ DSpark config # DSpark uses the DFlash-like target-hidden feedback path, then applies # a sequential Markov/confidence head outside the draft backbone engine. dspark_base: bool = False # When True, DSpark base export exposes DDTree parent/depth metadata. dspark_tree_base: bool = False is_dspark_draft_flag: bool = False dspark_target_layer_ids: List[int] = field(default_factory=list) dspark_block_size: int = 7 dspark_mask_token_id: int = 151669 dspark_enable_confidence_head: bool = False dspark_sample_from_anchor: bool = True dspark_confidence_head_with_markov: bool = False dspark_markov_head_type: str = "" dspark_markov_rank: int = 0 dspark_fc_native_precision: bool = False dspark_causal_proposal: bool = False # ------------------------------------------ sparse MoE config (Qwen3-style) # num_experts=0 means dense (no MoE) for Qwen/Mixtral-style keys; Nemotron-H instead reports # its expert count via n_routed_experts, so n_routed_experts > 0 also indicates MoE. num_experts: int = 0 n_routed_experts: int = 0 num_experts_per_tok: int = 0 # Expert MLP intermediate size (may differ from dense intermediate_size). moe_intermediate_size: int = 0 moe_shared_expert_intermediate_size: int = 0 # Optional latent dimension used by Nemotron-H routed experts. moe_latent_size: Optional[int] = None routed_scaling_factor: float = 1.0 n_group: int = 1 topk_group: int = 1 # MoE layer frequency: layer (i+1) % decoder_sparse_step == 0 → MoE. decoder_sparse_step: int = 1 # Layer indices that are always dense MLP (overrides decoder_sparse_step). mlp_only_layers: List[int] = field(default_factory=list) # Normalise top-k routing weights to sum to 1. # Note: the C++ Int4MoePlugin hardcodes renormalize=true, so this field # currently serves as documentation of the HF config value. norm_topk_prob: bool = True # Runtime vocabulary reduction. ``vocab_size`` remains the original # tokenizer/embedding size; this field is only the exported logits size. reduced_vocab_size: Optional[int] = None # ------------------------------------------ per-layer embeddings (Gemma4) # When > 0, Gemma4 E-model PLE is enabled. The ONNX graph receives # one runtime-provided ple_token_embeds_* tensor per decoder layer and # combines it with the context-aware projection from inputs_embeds. hidden_size_per_layer_input: int = 0 # Gemma4 E-model vocabulary for the runtime-side token-identity PLE table. # The table is exported as ple_embedding.safetensors and gathered by C++. vocab_size_per_layer_input: int = 0 # Gemma4: number of KV-shared layers (counted from the last layer backward). # Layers [num_hidden_layers - num_kv_shared_layers, num_hidden_layers) are # "KV-shared" — in HF they reuse KV states from the last non-shared layer # of the same attention type. Our export gives them independent K/V, but # they may have a 2× wider MLP (use_double_wide_mlp). num_kv_shared_layers: int = 0 # When True, KV-shared layers use 2× intermediate_size for their MLP. use_double_wide_mlp: bool = False # Gemma4 26B-A4B MoE: when True, each decoder layer has a routed MoE # block in addition to the dense MLP (parallel experts + router). enable_moe_block: bool = False # ------------------------------------------ tensor parallel # ``mapping`` is the single source of truth for parallel placement. # tp_size>1 returns a per-rank ONNX graph with col/row-parallel projections. mapping: Mapping = field(default_factory=Mapping) def __post_init__(self) -> None: # For standalone models, num_kv_shared_layers < num_hidden_layers is # required so that at least one non-shared donor layer exists. # Gemma4 assistant models (shares_target_kv=True) share KV from the # *target* model, so num_kv_shared_layers == num_hidden_layers is valid # — the donor-index logic is skipped at runtime for those. if (self.num_kv_shared_layers > 0 and self.num_kv_shared_layers >= self.num_hidden_layers and not self.shares_target_kv): raise ValueError( f"num_kv_shared_layers ({self.num_kv_shared_layers}) must be " f"less than num_hidden_layers ({self.num_hidden_layers})") # Resolve and validate the model family's attention scale. self.default_attention_scale = _get_attention_scaling( {}, self.head_dim, self.default_attention_scale) scale_config = ({} if self.attention_scaling is None else { "attention_scaling": self.attention_scaling }) self.attention_scaling = _get_attention_scaling( scale_config, self.head_dim, self.default_attention_scale) # ------------------------------------------------------------------ # Derived properties # ------------------------------------------------------------------ @property def tp_size(self) -> int: return self.mapping.tp_size @property def tp_rank(self) -> int: return self.mapping.tp_rank @property def is_eagle3_draft(self) -> bool: return self.draft_vocab_size is not None or self.is_eagle3_draft_flag @property def is_mtp_draft(self) -> bool: """True for a derived MTP draft config built from a base checkpoint.""" return bool(self.mtp_num_hidden_layers is not None and self.gdn_cfg is None and self.mamba_cfg is None and not self.mtp_base and not self.is_eagle3_draft and not self.is_dflash_draft and not self.is_jetspec_draft and not self.is_dspark_draft) @property def is_gemma4_mtp_draft(self) -> bool: """True for a paired Gemma4 assistant draft checkpoint.""" return self.gemma4_mtp_draft @property def is_diffusion_gemma(self) -> bool: return self.diffusion is not None or _is_diffusion_gemma_model_type( self.model_type) @property def is_dflash_draft(self) -> bool: return self.is_dflash_draft_flag @property def is_jetspec_draft(self) -> bool: return self.is_jetspec_draft_flag @property def is_dspark_draft(self) -> bool: return self.is_dspark_draft_flag @property def ple_enabled(self) -> bool: """True when Gemma4 per-layer embeddings are enabled.""" return self.hidden_size_per_layer_input > 0 @property def eagle3_target_hidden_size(self) -> int: return self.target_hidden_size or self.hidden_size @property def eagle3_num_target_layers(self) -> int: return len(self.eagle3_target_layer_ids) or 3 @property def is_hybrid(self) -> bool: return self.mamba_cfg is not None or self.gdn_cfg is not None @property def is_nemotron_h(self) -> bool: return (self.model_type or "").lower().startswith("nemotron_h") @property def num_attn_layers(self) -> int: """Total attention layers (includes Gemma4 sliding/full variants). Note: layers may have heterogeneous head dims — use per-layer configs (kv_layer_configs) for allocation, not this count alone. """ return sum(1 for t in self.layer_types if t in (LAYER_ATTN, "sliding_attention", "full_attention")) @property def num_mamba_layers(self) -> int: return sum(1 for t in self.layer_types if t == LAYER_MAMBA) @property def num_gdn_layers(self) -> int: return sum(1 for t in self.layer_types if t == LAYER_GDN) @property def num_mlp_layers(self) -> int: return sum(1 for t in self.layer_types if t == LAYER_MLP) @property def num_moe_layers(self) -> int: return sum(1 for t in self.layer_types if t == LAYER_MOE) @property def use_dual_rope(self) -> bool: return (self.sliding_rope_config is not None and self.full_rope_config is not None) @property def compute_dtype(self) -> "torch.dtype": # noqa: F821 import torch _MAP = { "bfloat16": torch.bfloat16, "float16": torch.float16, "float32": torch.float32, } return _MAP.get(self.torch_dtype, torch.bfloat16)
[docs] def for_rank(self, rank: int, world: int) -> "ModelConfig": """Return a per-rank copy of this config for TP. Divides head and intermediate sizes by *world* so each rank's model carries per-rank shapes. Usage: cfg = ModelConfig.from_pretrained( path, lambda head_dim: 1.0 / (float(head_dim)**0.5) ).for_rank(rank, world) model = CausalLM(cfg) load_weights(model, path, mapping=cfg.mapping) """ import copy if world == 1: return self parallel_dimensions = [ ("num_attention_heads", self.num_attention_heads), ("num_key_value_heads", self.num_key_value_heads), ("intermediate_size", self.intermediate_size), ] if self.gdn_cfg is not None: parallel_dimensions.extend(( ("gdn_cfg.num_key_heads", self.gdn_cfg.num_key_heads), ("gdn_cfg.num_value_heads", self.gdn_cfg.num_value_heads), )) for name, v in parallel_dimensions: if v % world: raise ValueError( f"TP world={world}: {name}={v} is not divisible by {world}" ) c = copy.deepcopy(self) c.mapping = Mapping(world_size=world, rank=rank, tp_size=world, tp_rank=rank) c.num_attention_heads //= world c.num_key_value_heads //= world c.intermediate_size //= world if c.gdn_cfg is not None: c.gdn_cfg.num_key_heads //= world c.gdn_cfg.num_value_heads //= world return c
# ------------------------------------------------------------------ # Factory # ------------------------------------------------------------------
[docs] @classmethod def from_pretrained( cls, model_dir: str, default_attention_scale: Callable[[int], float]) -> "ModelConfig": """Load a ModelConfig from a checkpoint directory. Loads architecture hyper-parameters via ``AutoConfig`` (see :func:`checkpoint_utils.load_checkpoint_config_dicts`) and then either ``hf_quant_config.json`` or the embedded ``quantization_config`` block to determine the quantisation scheme. ``default_attention_scale`` is a required model-family callable accepting ``head_dim``. ``has_qk_norm`` is auto-detected by scanning the safetensors key index for ``.q_norm.weight`` entries; the HunYuan ``.query_layernorm.weight`` spelling is honored only for model types whose loader remaps it (see :func:`_detect_has_qk_norm`). """ root, llm_dict = load_checkpoint_config_dicts(model_dir) root_model_type = root.get("model_type", "") model_type = llm_dict.get("model_type", "llama") is_diffusion_gemma = _is_diffusion_gemma_config(root, llm_dict) is_gemma4_assistant = _is_gemma4_assistant_model_type(root_model_type) if is_gemma4_assistant: model_type = "gemma4_assistant" elif is_diffusion_gemma: model_type = "diffusion_gemma" hidden_size = llm_dict["hidden_size"] num_attn_heads = llm_dict["num_attention_heads"] _check_num_attention_heads(num_attn_heads) head_dim = llm_dict.get("head_dim", hidden_size // num_attn_heads) global_head_dim = int(llm_dict.get("global_head_dim", 0) or 0) num_global_kv_heads = int( llm_dict.get("num_global_key_value_heads", 0) or 0) # Omni roots keep one ignore list for thinker + talker; scope it to the # sub-model this config represents (talker/CP go through their own # staged sub-config instead). submodel_prefix = "" if root.get("thinker_config") is not None and llm_dict is not root: submodel_prefix = "thinker." quant_dict = llm_dict if ("quantization_config" not in quant_dict and root.get("quantization_config") is not None): quant_dict = dict(llm_dict) quant_dict["quantization_config"] = root["quantization_config"] quant = _parse_quant(model_dir, quant_dict, submodel_prefix) raw_layer_types = _parse_raw_layer_types(llm_dict) layer_types = _parse_layer_types(llm_dict) attention_layer_types = _parse_attention_layer_types( llm_dict, llm_dict["num_hidden_layers"], model_type) num_kv_heads = llm_dict.get("num_key_value_heads", num_attn_heads) full_attention_only = (_is_gemma4_model_type(model_type) and attention_layer_types and all( layer_type == "full_attention" for layer_type in attention_layer_types)) if full_attention_only: if global_head_dim: head_dim = global_head_dim if num_global_kv_heads: num_kv_heads = num_global_kv_heads dual_rope_configs = _get_dual_rope_configs(llm_dict) mamba_cfg = _parse_mamba_cfg(llm_dict, layer_types, model_dir=model_dir) gdn_cfg = _parse_gdn_cfg(llm_dict, layer_types) has_qk_norm = _detect_has_qk_norm(model_dir, model_type) has_value_norm = _get_has_value_norm(llm_dict, model_type) default_attention_scale_value = float( default_attention_scale(head_dim)) attention_scaling = _get_attention_scaling( llm_dict, head_dim, default_attention_scale_value) embedding_scale = _get_embedding_scale(llm_dict, model_type, hidden_size) if _is_muse_glimmer_model_type(model_type): qk_scale_factor = llm_dict.get("qk_scale_factor") if qk_scale_factor is not None: attention_scaling = float(qk_scale_factor) / math.sqrt( head_dim) # Full-attention layers use NoPE; sliding layers use regular RoPE. if not dual_rope_configs: _mp = int(llm_dict.get("max_position_embeddings", 4096)) _sliding = { "rope_theta": _get_rope_theta(llm_dict), "rope_scaling": None, "partial_rotary_factor": 1.0, "max_position_embeddings": _mp, } _full = dict(_sliding) _full["rope_scaling"] = {"rope_type": "nope", "type": "nope"} dual_rope_configs = { "sliding_rope_config": _sliding, "full_rope_config": _full, } generation_config_path = os.path.join(model_dir, "generation_config.json") generation_config: Dict[str, Any] = {} if os.path.isfile(generation_config_path): with open(generation_config_path) as f: generation_config = json.load(f) diffusion_source_config = dict(root) diffusion_source_config.update(llm_dict) diffusion_config = (DiffusionConfig.from_hf(diffusion_source_config, generation_config) if is_diffusion_gemma else None) self_conditioning_size = 0 # MTP config mtp_num_hidden_layers = llm_dict.get("mtp_num_hidden_layers") if mtp_num_hidden_layers is not None: mtp_num_hidden_layers = int(mtp_num_hidden_layers) mtp_hybrid_override_pattern = llm_dict.get( "mtp_hybrid_override_pattern") mtp_layer_types = _parse_mtp_layer_types(llm_dict) num_nextn_predict_layers = int( llm_dict.get("num_nextn_predict_layers", 0) or 0) if (mtp_num_hidden_layers is None and num_nextn_predict_layers > 0 and mtp_layer_types): mtp_num_hidden_layers = len(mtp_layer_types) mtp_use_dedicated_embeddings = bool( llm_dict.get("mtp_use_dedicated_embeddings", False)) _validate_mtp_constraints( model_type=model_type, mtp_num_hidden_layers=mtp_num_hidden_layers, mtp_use_dedicated_embeddings=mtp_use_dedicated_embeddings, ) # EAGLE3 draft model fields architectures = llm_dict.get("architectures", []) or [] is_eagle3_draft_flag = any("Eagle3" in str(arch) for arch in architectures) or ("ttt_length" in llm_dict) draft_vocab_size = llm_dict.get("draft_vocab_size", None) target_hidden_size = llm_dict.get("target_hidden_size", None) eagle3_target_layer_ids = list( llm_dict.get("target_layer_ids", []) or []) if is_eagle3_draft_flag else [] if is_eagle3_draft_flag: draft_vocab_size = draft_vocab_size or llm_dict.get("vocab_size") target_hidden_size = target_hidden_size or hidden_size # Sliding window: active when use_sliding_window=True, or when # layer_types contains "sliding_attention" (Gemma4 convention). use_sw = llm_dict.get("use_sliding_window", False) or any( layer_type == "sliding_attention" for layer_type in raw_layer_types) sw_raw = llm_dict.get("sliding_window") if use_sw else None sliding_window_size = int(sw_raw) if sw_raw is not None else -1 skip_softmax_scale_factor = float( llm_dict.get("skip_softmax_scale_factor", 0.0)) use_vision_bidirectional_attention = bool( model_type in ("gemma4_unified", "gemma4_unified_text") and llm_dict.get("use_bidirectional_attention") == "vision") jetspec_config = (llm_dict.get("jetspec_config") or llm_dict.get("dflash_config") or {}) jetspec_target_layer_ids = list( jetspec_config.get("target_layer_ids", llm_dict.get("target_layer_ids", [])) or []) jetspec_block_size = int( jetspec_config.get("block_size", llm_dict.get("block_size", 16))) jetspec_mask_token_id = int( jetspec_config.get("mask_token_id", llm_dict.get("mask_token_id", 151669))) jetspec_causal_head = bool( jetspec_config.get("causal_head", llm_dict.get("causal_head", True))) # Sparse MoE fields. HF uses "num_local_experts" as the internal key # and maps "num_experts" → "num_local_experts" via attribute_map. num_experts = int( llm_dict.get("num_experts", llm_dict.get("num_local_experts", 0)) or 0) # Nemotron-H declares routed experts as ``n_routed_experts`` rather # than ``num_experts`` / ``num_local_experts``; treat either as MoE so # the sizes below are parsed instead of defaulting to 0. n_routed_experts = int(llm_dict.get("n_routed_experts", 0) or 0) if num_experts > 0 or n_routed_experts > 0: num_experts_per_tok = int( llm_dict.get("num_experts_per_tok", llm_dict.get("top_k_experts", 0)) or 0) moe_intermediate_size = int( llm_dict.get("moe_intermediate_size", 0) or 0) else: num_experts_per_tok = 0 moe_intermediate_size = 0 # HF Qwen3-Omni MoE Talker uses the un-prefixed name; HF NemotronH / # other MoE families use ``moe_shared_expert_intermediate_size``. moe_shared_expert_intermediate_size = int( llm_dict.get("moe_shared_expert_intermediate_size", llm_dict.get("shared_expert_intermediate_size", 0)) or 0) routed_scaling_factor = float( llm_dict.get("routed_scaling_factor", 1.0)) n_group = int(llm_dict.get("n_group", 1)) topk_group = int(llm_dict.get("topk_group", 1)) decoder_sparse_step = int(llm_dict.get("decoder_sparse_step", 1)) mlp_only_layers = list(llm_dict.get("mlp_only_layers") or []) norm_topk_prob = bool(llm_dict.get("norm_topk_prob", True)) moe_latent_size = llm_dict.get("moe_latent_size", None) if moe_latent_size is not None: moe_latent_size = int(moe_latent_size) intermediate_size = int( llm_dict.get("intermediate_size") or llm_dict.get("shared_expert_intermediate_size") or llm_dict.get("moe_shared_expert_intermediate_size") or llm_dict.get("moe_intermediate_size", 0)) if self_conditioning_size <= 0: self_conditioning_size = int( llm_dict.get( "self_conditioning_size", root.get("self_conditioning_size", intermediate_size)) or intermediate_size) return cls( model_type=model_type, hidden_size=hidden_size, num_hidden_layers=llm_dict["num_hidden_layers"], num_attention_heads=num_attn_heads, num_key_value_heads=num_kv_heads, intermediate_size=intermediate_size, head_dim=head_dim, global_head_dim=global_head_dim, num_global_key_value_heads=num_global_kv_heads, rms_norm_eps=_get_rms_norm_eps(llm_dict, model_type), vocab_size=llm_dict["vocab_size"], rope_theta=_get_rope_theta(llm_dict), max_position_embeddings=llm_dict.get("max_position_embeddings", 4096), default_attention_scale=default_attention_scale_value, rope_scaling=_select_rope_scaling(llm_dict), original_max_position_embeddings=llm_dict.get( "original_max_position_embeddings", None), partial_rotary_factor=_get_partial_rotary_factor(llm_dict), hidden_activation=llm_dict.get("hidden_activation", llm_dict.get("hidden_act", "silu")), num_code_groups=int(llm_dict.get("num_code_groups", 0) or 0), sliding_rope_config=dual_rope_configs.get("sliding_rope_config"), full_rope_config=dual_rope_configs.get("full_rope_config"), has_qk_norm=has_qk_norm, qk_norm_post_rope=(model_type in _QK_NORM_POST_ROPE_MODEL_TYPES), has_value_norm=has_value_norm, attention_bias=bool(llm_dict.get("attention_bias", False)), attention_scaling=attention_scaling, attention_k_eq_v=(True if is_diffusion_gemma else bool( llm_dict.get("attention_k_eq_v", False))), encoder_layer_scalars=list( llm_dict.get("encoder_layer_scalars", root.get("encoder_layer_scalars", [])) or []), decoder_layer_scalars=list( llm_dict.get("decoder_layer_scalars", root.get("decoder_layer_scalars", [])) or []), self_conditioning_size=self_conditioning_size, diffusion=diffusion_config, embedding_scale=embedding_scale, final_logit_softcapping=llm_dict.get("final_logit_softcapping", None), output_multiplier=float(llm_dict.get("output_multiplier", 1.0)), post_norm_eps=(float(llm_dict["post_norm_eps"]) if llm_dict.get("post_norm_eps") is not None else None), torch_dtype=llm_dict.get("torch_dtype", llm_dict.get("dtype", "bfloat16")), tie_word_embeddings=llm_dict.get("tie_word_embeddings", False), sliding_window_size=sliding_window_size, skip_softmax_scale_factor=skip_softmax_scale_factor, use_vision_bidirectional_attention= use_vision_bidirectional_attention, layer_types=layer_types, attention_layer_types=attention_layer_types, quant=quant, mamba_cfg=mamba_cfg, gdn_cfg=gdn_cfg, attn_output_gate=bool(llm_dict.get("attn_output_gate", False)), mtp_num_hidden_layers=mtp_num_hidden_layers, mtp_use_dedicated_embeddings=mtp_use_dedicated_embeddings, mtp_hybrid_override_pattern=mtp_hybrid_override_pattern, mtp_layer_types=mtp_layer_types, mtp_base=bool(llm_dict.get("mtp_base", False)), root_model_type=root_model_type, raw_layer_types=raw_layer_types, rope_parameters=llm_dict.get("rope_parameters", None), backbone_hidden_size=int( llm_dict.get("backbone_hidden_size", 0) or 0), assistant_hidden_size=(hidden_size if is_gemma4_assistant else 0), shares_target_kv=is_gemma4_assistant, has_own_kv_cache=not is_gemma4_assistant, constant_draft_positions=is_gemma4_assistant, returns_feedback_hidden=is_gemma4_assistant, use_ordered_embeddings=bool( llm_dict.get("use_ordered_embeddings", False)), num_centroids=int(llm_dict.get("num_centroids", 0) or 0), centroid_intermediate_top_k=int( llm_dict.get("centroid_intermediate_top_k", 0) or 0), mtp_tree_base=bool(llm_dict.get("mtp_tree_base", False)), dflash_base=bool(llm_dict.get("dflash_base", False)), dflash_tree_base=bool(llm_dict.get("dflash_tree_base", False)), jetspec_base=bool(llm_dict.get("jetspec_base", False)), jetspec_tree_base=bool(llm_dict.get("jetspec_tree_base", False)), jetspec_target_layer_ids=jetspec_target_layer_ids, jetspec_block_size=jetspec_block_size, jetspec_mask_token_id=jetspec_mask_token_id, jetspec_causal_head=jetspec_causal_head, dflash_target_layer_ids=list( (llm_dict.get("dflash_config", {}) or {}).get("target_layer_ids") or llm_dict.get("eagle_aux_hidden_state_layer_ids") or []), dspark_base=bool(llm_dict.get("dspark_base", False)), dspark_tree_base=bool(llm_dict.get("dspark_tree_base", False)), num_deepstack_features=_parse_num_deepstack_features( llm_dict, model_type, root_config=root), accept_hidden_layer=_parse_accept_hidden_layer(llm_dict, root_config=root), draft_vocab_size=draft_vocab_size, target_hidden_size=target_hidden_size, is_eagle3_draft_flag=is_eagle3_draft_flag, eagle3_target_layer_ids=eagle3_target_layer_ids, num_experts=num_experts, n_routed_experts=n_routed_experts, num_experts_per_tok=num_experts_per_tok, moe_intermediate_size=moe_intermediate_size, moe_shared_expert_intermediate_size= moe_shared_expert_intermediate_size, moe_latent_size=moe_latent_size, routed_scaling_factor=routed_scaling_factor, n_group=n_group, topk_group=topk_group, decoder_sparse_step=decoder_sparse_step, mlp_only_layers=mlp_only_layers, norm_topk_prob=norm_topk_prob, hidden_size_per_layer_input=int( llm_dict.get("hidden_size_per_layer_input", 0) or 0), vocab_size_per_layer_input=int( llm_dict.get("vocab_size_per_layer_input", 0) or 0), num_kv_shared_layers=int( llm_dict.get("num_kv_shared_layers", 0) or 0), use_double_wide_mlp=bool(llm_dict.get("use_double_wide_mlp", False)), enable_moe_block=bool( llm_dict.get( "enable_moe_block", model_type in ( "gemma4", "gemma4_text", "gemma4_unified", "gemma4_unified_text", "diffusion_gemma", ) and num_experts > 0 and moe_intermediate_size > 0)), )
# --------------------------------------------------------------------------- # Internal helpers # --------------------------------------------------------------------------- # When HF omits ``deepstack_visual_indexes`` / ``num_deepstack_features``, # these model_types are known to expect three visual deepstack injections # at runtime (Thinker side only). Listed explicitly — substring matching # (``"qwen3_omni" in "qwen3_omni_moe_talker"``) would otherwise mis-classify # Talker / CodePredictor configs as deepstack producers and bake 3 dangling # input ports into their engines. _DEEPSTACK_MODEL_TYPES = frozenset({ # HF root configs that wrap a deepstack-emitting visual encoder "qwen3_vl", "qwen3_omni", "qwen3_omni_moe", # Standalone Thinker text-LLM configs (after quant export); the Thinker # still consumes deepstack inputs from the separately exported visual # encoder at runtime. "qwen3_vl_text", "qwen3_omni_text", "qwen3_omni_moe_text", }) _QWEN3_5_MTP_CONFIG_MODEL_TYPES = frozenset({ "qwen3_5", "qwen3_5_text", "qwen3_5_moe", "qwen3_5_moe_text", "qwen3_omni_next_text", "qwen3_omni_next_text_moe", }) def make_mtp_draft_config(base_config: ModelConfig) -> ModelConfig: """Derive the currently supported MTP draft config from a base config.""" _validate_mtp_constraints( model_type=base_config.model_type, mtp_num_hidden_layers=base_config.mtp_num_hidden_layers, mtp_use_dedicated_embeddings=base_config.mtp_use_dedicated_embeddings, ) mtp_num_hidden_layers = base_config.mtp_num_hidden_layers if mtp_num_hidden_layers is None: raise ValueError( "MTP draft config requires mtp_num_hidden_layers in the base config." ) # Nemotron-H: Exclude all draft ``layers.*`` modules from quantization; # the untouched lm_head keeps the base quant type. if (base_config.model_type or "").lower().startswith("nemotron_h"): draft_layer_types = list(base_config.mtp_layer_types) if len(draft_layer_types) != mtp_num_hidden_layers: declared = (base_config.mtp_hybrid_override_pattern or base_config.mtp_layer_types) raise ValueError( f"MTP draft stack {declared!r} yields " f"{len(draft_layer_types)} layers != mtp_num_hidden_layers " f"{mtp_num_hidden_layers}") draft_quant = replace( base_config.quant, excluded=list(base_config.quant.excluded) + ["layers.*"], ) return replace( base_config, num_hidden_layers=mtp_num_hidden_layers, layer_types=draft_layer_types, mamba_cfg=None, gdn_cfg=None, mtp_base=False, quant=draft_quant, tie_word_embeddings=False, ) # The draft is quantized iff its FFN compute weights are quantized # (routed experts for MoE, dense MLP otherwise) — other modules may be # excluded in any recipe. Excluded entries are exact names or globs. ffn_probe = ("mtp.layers.0.mlp.experts.0.gate_proj" if base_config.num_experts > 0 else "mtp.layers.0.mlp.gate_proj") mtp_is_quantized = not any( fnmatch.fnmatch(ffn_probe, e) for e in base_config.quant.excluded) if mtp_is_quantized: # MTP draft modules are independently quantized. Strip base-model # layer_overrides that use module paths absent from the draft # (e.g. "model.layers.0.linear_attn.in_proj_qkv") — they would # cause spurious FP16 fallback for draft-specific layers like "fc". # Keep entries whose keys also exist in the draft module namespace # (e.g. "lm_head") so that lm_head_quantization is honoured. _DRAFT_MODULE_PREFIXES = ("lm_head", "fc", "layers.", "norm") draft_overrides = { k: v for k, v in base_config.quant.layer_overrides.items() if any(k == p or k.startswith(p) for p in _DRAFT_MODULE_PREFIXES) } draft_group_sizes = { k: v for k, v in base_config.quant.layer_group_sizes.items() if any(k == p or k.startswith(p) for p in _DRAFT_MODULE_PREFIXES) } # Preserve MTP-specific exclusions (e.g. mtp.lm_head when lm_head # is FP16) but drop base-model exclusions irrelevant to the draft. draft_excluded = [ e[len("mtp."):] for e in base_config.quant.excluded if e.startswith("mtp.") ] # is_mixed_precision=False so unlisted modules (fc, q_proj, etc.) # fall back to the dominant quant type, not FP16. Explicit # overrides (lm_head→fp8) still take effect via layer_overrides. draft_quant = replace(base_config.quant, excluded=draft_excluded, layer_overrides=draft_overrides, layer_group_sizes=draft_group_sizes, is_mixed_precision=False) else: # MTP draft lm_head is borrowed from the base model and may itself be quantized. lm_head_overrides = { k: v for k, v in base_config.quant.layer_overrides.items() if k == "lm_head" or k.startswith("lm_head.") } lm_head_group_sizes = { k: v for k, v in base_config.quant.layer_group_sizes.items() if k == "lm_head" or k.startswith("lm_head.") } draft_quant = QuantConfig(layer_overrides=lm_head_overrides, layer_group_sizes=lm_head_group_sizes) return replace( base_config, num_hidden_layers=mtp_num_hidden_layers, layer_types=[LAYER_ATTN] * mtp_num_hidden_layers, gdn_cfg=None, mtp_base=False, quant=draft_quant, tie_word_embeddings=False, ) def make_dspark_draft_config( draft_dir: str, default_attention_scale: Callable[[int], float]) -> ModelConfig: """Build a DSpark draft ModelConfig from a DeepSpec DSpark checkpoint.""" _, llm_dict = load_checkpoint_config_dicts(draft_dir) dspark_config = (llm_dict.get("dspark_config") or llm_dict.get("dflash_config", {}) or {}) quant = _parse_quant(draft_dir, llm_dict) target_layer_ids = list( dspark_config.get("target_layer_ids", llm_dict.get("target_layer_ids", []))) if not target_layer_ids: raise ValueError( "DSpark draft config requires target_layer_ids in config.json.") model_type = llm_dict.get("model_type", "qwen3") raw_layer_types = _parse_raw_layer_types(llm_dict) layer_types = _parse_layer_types(llm_dict) attention_layer_types = _parse_attention_layer_types( llm_dict, llm_dict["num_hidden_layers"], model_type) full_attention_only = (_is_gemma4_model_type(model_type) and attention_layer_types and all(layer_type == "full_attention" for layer_type in attention_layer_types)) head_dim = llm_dict.get( "global_head_dim" if full_attention_only else "head_dim", llm_dict.get( "head_dim", llm_dict["hidden_size"] // llm_dict["num_attention_heads"])) global_head_dim = int(llm_dict.get("global_head_dim", 0) or 0) num_global_kv_heads = int( llm_dict.get("num_global_key_value_heads", 0) or 0) num_kv_heads = ( num_global_kv_heads if full_attention_only and num_global_kv_heads else llm_dict.get("num_key_value_heads", llm_dict["num_attention_heads"])) dual_rope_configs = _get_dual_rope_configs(llm_dict) use_sw = llm_dict.get("use_sliding_window", False) or any( layer_type == "sliding_attention" for layer_type in raw_layer_types) sw_raw = llm_dict.get("sliding_window") if use_sw else None sliding_window_size = int(sw_raw) if sw_raw is not None else -1 default_attention_scale_value = float(default_attention_scale(head_dim)) default_mask_token_id = 4 if _is_gemma4_model_type(model_type) else 151669 return ModelConfig( model_type=model_type, hidden_size=llm_dict["hidden_size"], num_hidden_layers=llm_dict["num_hidden_layers"], num_attention_heads=llm_dict["num_attention_heads"], num_key_value_heads=num_kv_heads, intermediate_size=llm_dict["intermediate_size"], head_dim=head_dim, global_head_dim=global_head_dim, num_global_key_value_heads=num_global_kv_heads, rms_norm_eps=llm_dict.get("rms_norm_eps", 1e-6), vocab_size=llm_dict["vocab_size"], rope_theta=_get_rope_theta(llm_dict), max_position_embeddings=llm_dict.get("max_position_embeddings", 4096), default_attention_scale=default_attention_scale_value, rope_scaling=_select_rope_scaling(llm_dict), partial_rotary_factor=_get_partial_rotary_factor(llm_dict), hidden_activation=llm_dict.get("hidden_activation", llm_dict.get("hidden_act", "silu")), sliding_rope_config=dual_rope_configs.get("sliding_rope_config"), full_rope_config=dual_rope_configs.get("full_rope_config"), has_qk_norm=True, has_value_norm=_get_has_value_norm(llm_dict, model_type), attention_bias=bool(llm_dict.get("attention_bias", False)), attention_scaling=_get_attention_scaling( llm_dict, head_dim, default_attention_scale_value), attention_k_eq_v=bool(llm_dict.get("attention_k_eq_v", False)), attention_sink_bias=bool( dspark_config.get("attention_sink_bias", llm_dict.get("attention_sink_bias", False))), final_logit_softcapping=llm_dict.get("final_logit_softcapping", None), torch_dtype=llm_dict.get("torch_dtype", llm_dict.get("dtype", "bfloat16")), tie_word_embeddings=False, sliding_window_size=sliding_window_size, layer_types=layer_types, attention_layer_types=attention_layer_types, raw_layer_types=raw_layer_types, rope_parameters=llm_dict.get("rope_parameters", None), is_dspark_draft_flag=True, dspark_causal_proposal=bool( dspark_config.get( "causal", llm_dict.get("dflash_query_causal", (llm_dict.get("dflash_config") or {}).get("causal", False)))), dspark_target_layer_ids=target_layer_ids, dspark_block_size=int( dspark_config.get("block_size", llm_dict.get("block_size", 7))), dspark_mask_token_id=int( dspark_config.get( "mask_token_id", llm_dict.get("mask_token_id", default_mask_token_id))), dspark_enable_confidence_head=bool( dspark_config.get("enable_confidence_head", llm_dict.get("enable_confidence_head", False))), dspark_sample_from_anchor=bool( dspark_config.get("sample_from_anchor", llm_dict.get("sample_from_anchor", True))), dspark_confidence_head_with_markov=bool( dspark_config.get( "confidence_head_with_markov", llm_dict.get("confidence_head_with_markov", False))), dspark_markov_head_type=str( dspark_config.get("markov_head_type", llm_dict.get("markov_head_type", ""))), dspark_markov_rank=int( dspark_config.get("markov_rank", llm_dict.get("markov_rank", 0))), quant=quant, ) def make_dflash_draft_config( draft_dir: str, default_attention_scale: Callable[[int], float], target_vocab_size: Optional[int] = None) -> ModelConfig: """Build a DFlash draft ModelConfig from the draft checkpoint directory. Now quantization-aware: if the draft directory contains ``hf_quant_config.json`` (e.g. from DFlash draft NVFP4 quantization), the quant config is parsed so ``make_linear`` will produce the correct Linear class (``NVFP4Linear`` etc.) during ONNX export. """ _, llm_dict = load_checkpoint_config_dicts(draft_dir) dflash_config = llm_dict.get("dflash_config", {}) or {} # Parse quantization config from the draft checkpoint directory. # For FP16 draft checkpoints this returns QuantConfig() (no quant). quant = _parse_quant(draft_dir, llm_dict) # The fc feature projector must stay dense FP16 (the draft model asserts # this). Finalized NVFP4 drafts ship a packed fc; excluding it here keeps # ``make_linear`` producing FP16Linear, and the loader dequantizes the # packed checkpoint tensors into it (see ``model.py``). if "fc" not in quant.excluded: quant.excluded = list(quant.excluded) + ["fc"] model_type = llm_dict.get("model_type", "qwen3") _check_num_attention_heads(llm_dict["num_attention_heads"]) raw_layer_types = _parse_raw_layer_types(llm_dict) layer_types = _parse_layer_types(llm_dict) attention_layer_types = _parse_attention_layer_types( llm_dict, llm_dict["num_hidden_layers"], model_type) full_attention_only = (_is_gemma4_model_type(model_type) and attention_layer_types and all(layer_type == "full_attention" for layer_type in attention_layer_types)) head_dim = llm_dict.get( "global_head_dim" if full_attention_only else "head_dim", llm_dict.get( "head_dim", llm_dict["hidden_size"] // llm_dict["num_attention_heads"])) global_head_dim = int(llm_dict.get("global_head_dim", 0) or 0) num_global_kv_heads = int( llm_dict.get("num_global_key_value_heads", 0) or 0) num_kv_heads = ( num_global_kv_heads if full_attention_only and num_global_kv_heads else llm_dict.get("num_key_value_heads", llm_dict["num_attention_heads"])) dual_rope_configs = _get_dual_rope_configs(llm_dict) use_sw = llm_dict.get("use_sliding_window", False) or any( layer_type == "sliding_attention" for layer_type in raw_layer_types) sw_raw = llm_dict.get("sliding_window") if use_sw else None sliding_window_size = int(sw_raw) if sw_raw is not None else -1 default_attention_scale_value = float(default_attention_scale(head_dim)) default_mask_token_id = 4 if _is_gemma4_model_type(model_type) else 248070 vocab_size = llm_dict.get("vocab_size", target_vocab_size) if vocab_size is None: raise ValueError("DFlash draft config must provide vocab_size") return ModelConfig( model_type=model_type, hidden_size=llm_dict["hidden_size"], num_hidden_layers=llm_dict["num_hidden_layers"], num_attention_heads=llm_dict["num_attention_heads"], num_key_value_heads=num_kv_heads, intermediate_size=llm_dict["intermediate_size"], head_dim=head_dim, global_head_dim=global_head_dim, num_global_key_value_heads=num_global_kv_heads, rms_norm_eps=llm_dict.get("rms_norm_eps", 1e-6), vocab_size=int(vocab_size), rope_theta=_get_rope_theta(llm_dict), max_position_embeddings=llm_dict.get("max_position_embeddings", 4096), default_attention_scale=default_attention_scale_value, rope_scaling=_select_rope_scaling(llm_dict), partial_rotary_factor=_get_partial_rotary_factor(llm_dict), hidden_activation=llm_dict.get("hidden_activation", llm_dict.get("hidden_act", "silu")), sliding_rope_config=dual_rope_configs.get("sliding_rope_config"), full_rope_config=dual_rope_configs.get("full_rope_config"), has_qk_norm=True, has_value_norm=_get_has_value_norm(llm_dict, model_type), attention_scaling=_get_attention_scaling( llm_dict, head_dim, default_attention_scale_value), attention_k_eq_v=bool(llm_dict.get("attention_k_eq_v", False)), final_logit_softcapping=llm_dict.get("final_logit_softcapping", None), torch_dtype=llm_dict.get("torch_dtype", "bfloat16"), tie_word_embeddings=False, sliding_window_size=sliding_window_size, layer_types=layer_types, attention_layer_types=attention_layer_types, raw_layer_types=raw_layer_types, rope_parameters=llm_dict.get("rope_parameters", None), is_dflash_draft_flag=True, dflash_target_layer_ids=list( dflash_config.get( "target_layer_ids", llm_dict.get("target_layer_ids", [1, 8, 15, 22, 29]))), dflash_block_size=int( dflash_config.get("block_size", llm_dict.get("block_size", 16))), dflash_mask_token_id=int( dflash_config.get( "mask_token_id", llm_dict.get("mask_token_id", default_mask_token_id))), quant=quant, ) def make_dflash2_draft_config( draft_dir: str, default_attention_scale: Callable[[int], float]) -> ModelConfig: """Build and validate the production DFlash2 draft contract.""" root_dict, llm_dict = load_checkpoint_config_dicts(draft_dir) resolved = resolve_dflash_contract(root_dict, llm_dict) if resolved.version != DFlashVersion.V2: raise ValueError( "DFlash2 draft checkpoint requires architecture DFlash2DraftModel") config = make_dflash_draft_config(draft_dir, default_attention_scale) if config.num_hidden_layers != 5: raise ValueError( "DFlash2 production checkpoint requires exactly five draft layers") return replace( config, dflash_version=resolved.version, dflash_target_layer_ids=list(resolved.target_layer_ids), dflash_block_size=resolved.block_size, dflash_mask_token_id=resolved.mask_token_id, is_dflash_draft_flag=True, dflash2_target_layer_ids=list(resolved.target_layer_ids), dflash2_block_size=resolved.block_size, dflash2_mask_token_id=resolved.mask_token_id, dflash2_is_causal=resolved.is_causal, dflash2_conv_kernel_size=resolved.conv_kernel_size, dflash2_conv_group_size=resolved.conv_group_size, dflash2_selector_rank=resolved.selector_rank, dflash2_selector_top_k=resolved.selector_top_k, ) def make_jetspec_draft_config( draft_dir: str, default_attention_scale: Callable[[int], float]) -> ModelConfig: """Build a JetSpec draft config from an official JetSpec checkpoint. The public JetSpec Qwen3 checkpoint stores its metadata in ``dflash_config`` because JetSpec reuses the DFlash draft-head implementation. Persist the exported runtime config under ``jetspec_config`` while also filling the DFlash fields consumed by :class:`DFlashDraftModel`. """ _, llm_dict = load_checkpoint_config_dicts(draft_dir) base = make_dflash_draft_config(draft_dir, default_attention_scale) jetspec_config = (llm_dict.get("jetspec_config") or llm_dict.get("dflash_config") or {}) target_layer_ids = list( jetspec_config.get("target_layer_ids", llm_dict.get("target_layer_ids", [])) or []) if not target_layer_ids: raise ValueError( "JetSpec draft config requires target_layer_ids in config.json.") block_size = int( jetspec_config.get("block_size", llm_dict.get("block_size", base.dflash_block_size))) mask_token_id = int( jetspec_config.get("mask_token_id", llm_dict.get("mask_token_id", 151669))) causal_head = bool( jetspec_config.get("causal_head", llm_dict.get("causal_head", True))) if not causal_head: raise ValueError( "JetSpec draft config requires causal_head=true; use DFlash for non-causal block drafts." ) return replace( base, is_dflash_draft_flag=False, is_jetspec_draft_flag=True, jetspec_target_layer_ids=target_layer_ids, jetspec_block_size=block_size, jetspec_mask_token_id=mask_token_id, jetspec_causal_head=causal_head, dflash_target_layer_ids=target_layer_ids, dflash_block_size=block_size, dflash_mask_token_id=mask_token_id, ) def _parse_num_deepstack_features( config: dict, model_type: str, *, root_config: Optional[Dict[str, Any]] = None, ) -> int: """Return the number of deepstack visual features for ONNX / runtime. 1. Explicit ``num_deepstack_features`` on the LLM (text) dict, then on the root dict. 2. ``len(vision_config["deepstack_visual_indexes"])`` on the root config. 3. Fallback: ``3`` for ``qwen3_vl`` / ``qwen3_omni`` when still unknown, else ``0``. """ raw = config.get("num_deepstack_features") if raw is not None: return int(raw) if root_config is not None: raw_root = root_config.get("num_deepstack_features") if raw_root is not None: return int(raw_root) vision = root_config.get("vision_config") if isinstance(vision, dict): indexes = vision.get("deepstack_visual_indexes") if isinstance(indexes, (list, tuple)) and len(indexes) > 0: return len(indexes) root_mt = (root_config or {}).get("model_type") or "" # Strict equality, not substring — ``"qwen3_omni" in "qwen3_omni_next"`` # is True and would otherwise report 3 deepstack features for Qwen3-Next Omni # (whose ``deepstack_visual_indexes`` is empty -> 0 features). if model_type in _DEEPSTACK_MODEL_TYPES or root_mt in _DEEPSTACK_MODEL_TYPES: return 3 return 0 def _parse_accept_hidden_layer( config: dict, *, root_config: Optional[Dict[str, Any]] = None, ) -> int: """Return ``accept_hidden_layer`` (Thinker decoder layer count whose output is handed off to the Qwen3-Omni Talker), or -1 when not applicable. Lookup order: 1. ``accept_hidden_layer`` at the LLM (text) dict top level. This is the layout written by the standalone-Thinker quant export, which promotes the field out of ``talker_config`` into the Thinker root. 2. ``root_config["talker_config"]["accept_hidden_layer"]`` for the full multimodal HF config. 3. ``root_config["accept_hidden_layer"]`` as a defensive fallback for standalone Talker checkpoints where this field already lives at root. """ raw = config.get("accept_hidden_layer") if raw is not None: return int(raw) if root_config is not None: talker_cfg = root_config.get("talker_config") if isinstance(talker_cfg, dict): raw = talker_cfg.get("accept_hidden_layer") if raw is not None: return int(raw) raw = root_config.get("accept_hidden_layer") if raw is not None: return int(raw) return -1 def _validate_mtp_constraints( *, model_type: str, mtp_num_hidden_layers: Optional[int], mtp_use_dedicated_embeddings: bool, ) -> None: """Validate the currently supported MTP config subset.""" if mtp_num_hidden_layers is None and not mtp_use_dedicated_embeddings: return if (model_type or "").lower().startswith("nemotron_h"): if mtp_use_dedicated_embeddings: raise NotImplementedError( "Dedicated MTP embeddings are not supported for Nemotron-H MTP." ) return if model_type not in _QWEN3_5_MTP_CONFIG_MODEL_TYPES: raise NotImplementedError( "MTP config parsing is only supported for Qwen3.5 and Nemotron-H " "checkpoints.") if mtp_num_hidden_layers != 1: raise NotImplementedError( "Only mtp_num_hidden_layers == 1 is supported for Qwen3.5 MTP.") if mtp_use_dedicated_embeddings: raise NotImplementedError( "Dedicated MTP embeddings are not supported for Qwen3.5 MTP.") def _parse_raw_layer_types(config: dict) -> List[str]: """Return checkpoint layer type strings without canonicalization.""" raw = config.get("layers_block_type") or config.get("layer_types") if raw is None: return [] return [str(layer_type) for layer_type in raw] def _is_nemotron_h_config(config: dict) -> bool: return (config.get("model_type") or "").lower().startswith("nemotron_h") def _canonical_layer_type(block_type: str, is_nemotron_h: bool) -> str: """Map one checkpoint block-type name onto a canonical layer label. ``"linear_attention"`` is ambiguous: it covers any sub-quadratic mixer, so it denotes Mamba2 for NemotronH and GatedDeltaNet for Qwen3.5. The model family therefore resolves it, along with ``"full_attention"``, which Gemma4 keeps verbatim for per-layer head_dim dispatch. """ bt = str(block_type).lower() if bt == "linear_attention": return LAYER_MAMBA if is_nemotron_h else LAYER_GDN if "mamba" in bt: return LAYER_MAMBA if bt == "moe": return LAYER_MOE if "mlp" in bt: return LAYER_MLP if bt in _VALID_ATTENTION_LAYER_TYPES: return LAYER_ATTN if is_nemotron_h else bt return LAYER_ATTN def _parse_mtp_layer_types(config: dict) -> List[str]: """Return the MTP draft stack's per-layer block types. NemotronH declares the stack either as ``mtp_layers_block_type`` (a list, which transformers >= 5.14 rewrites into the ``linear_attention`` / ``full_attention`` spelling) or as the legacy ``mtp_hybrid_override_pattern`` string (e.g. ``"*E"``). """ raw = config.get("mtp_layers_block_type") if raw: is_nemotron_h = _is_nemotron_h_config(config) return [_canonical_layer_type(bt, is_nemotron_h) for bt in raw] pattern = config.get("mtp_hybrid_override_pattern") or "" return [ _HYBRID_PATTERN_MAP[ch] for ch in pattern if ch in _HYBRID_PATTERN_MAP ] def _parse_layer_types(config: dict) -> List[str]: """Return per-layer block type list from config. Reads ``layers_block_type`` or ``layer_types`` directly if present, each entry resolved by :func:`_canonical_layer_type`. For models using ``hybrid_override_pattern`` (e.g. NemotronH), parses the pattern string where ``M`` = mamba, ``-`` = mlp, ``*`` = attention. Falls back to all attention layers. """ raw = config.get("layers_block_type") or config.get("layer_types") if raw is not None: is_nemotron_h = _is_nemotron_h_config(config) return [_canonical_layer_type(bt, is_nemotron_h) for bt in raw] pattern = config.get("hybrid_override_pattern") if pattern is not None: return [ _HYBRID_PATTERN_MAP[ch] for ch in pattern if ch in _HYBRID_PATTERN_MAP ] n = config["num_hidden_layers"] return [LAYER_ATTN] * n def _parse_mamba_cfg(config: dict, layer_types: List[str], model_dir: str = "") -> Optional[MambaConfig]: """Return a MambaConfig if any layer is a mamba layer, else None. ``model_dir`` is used to detect ``conv_dim`` from the actual checkpoint weight shape when the config does not declare it explicitly. This resolves a circular dependency: ``conv_dim`` depends on ``n_groups`` and vice-versa. """ if LAYER_MAMBA not in layer_types: return None num_heads = config.get("mamba_num_heads", 0) head_dim = config.get("mamba_head_dim", 0) ssm_state_size = config.get("ssm_state_size", 0) conv_kernel = config.get("conv_kernel", config.get("mamba_d_conv", 4)) d_inner = num_heads * head_dim n_groups = config.get( "n_groups", config.get("mamba_n_groups", config.get("mamba_num_groups", 1))) if "conv_dim" in config: conv_dim = config["conv_dim"] else: # Try to read conv_dim from the actual checkpoint shape to avoid # circular dependency with n_groups when neither key is in config. detected = _detect_mamba_conv_dim(model_dir) if model_dir else 0 if detected > 0: conv_dim = detected else: conv_dim = d_inner + 2 * n_groups * ssm_state_size # Back-derive n_groups from conv_dim if not explicit in config if n_groups == 1 and conv_dim > d_inner: derived = (conv_dim - d_inner) // (2 * ssm_state_size) n_groups = derived if derived > 0 else 1 return MambaConfig( num_heads=num_heads, head_dim=head_dim, ssm_state_size=ssm_state_size, conv_dim=conv_dim, conv_kernel=conv_kernel, n_groups=n_groups, ) def _detect_mamba_conv_dim(model_dir: str) -> int: """Return conv_dim by reading the first Mamba conv1d.weight shape from the checkpoint. The conv1d weight has shape ``[conv_dim, 1, conv_kernel]`` so ``shape[0]`` gives ``conv_dim`` directly — no tensor data is loaded. Returns 0 if not found (caller falls back to formula-based derivation). """ try: index_path = os.path.join(model_dir, "model.safetensors.index.json") if os.path.exists(index_path): with open(index_path) as f: index = json.load(f) weight_map: dict = index.get("weight_map", {}) shard_for_key: Optional[str] = None target_key: Optional[str] = None for k, shard in weight_map.items(): if k.endswith(".mixer.conv1d.weight"): shard_for_key = shard target_key = k break if shard_for_key and target_key: shard_path = os.path.join(model_dir, shard_for_key) from safetensors import safe_open with safe_open(shard_path, framework="pt") as f: return f.get_slice(target_key).get_shape()[0] else: single_path = os.path.join(model_dir, "model.safetensors") if os.path.exists(single_path): from safetensors import safe_open with safe_open(single_path, framework="pt") as f: for k in f.keys(): if k.endswith(".mixer.conv1d.weight"): return f.get_slice(k).get_shape()[0] except (OSError, KeyError, ImportError): pass return 0 def _parse_gdn_cfg(config: dict, layer_types: List[str]) -> Optional[GdnConfig]: """Return a GdnConfig if any layer is a GDN (linear_attention) layer, else None.""" if LAYER_GDN not in layer_types: return None return GdnConfig( num_key_heads=config.get("linear_num_key_heads", 0), num_value_heads=config.get("linear_num_value_heads", 0), key_head_dim=config.get("linear_key_head_dim", 0), value_head_dim=config.get("linear_value_head_dim", 0), conv_kernel=config.get("linear_conv_kernel_dim", 4), ) def _get_partial_rotary_factor(llm_dict: Dict[str, Any]) -> float: """Extract partial_rotary_factor from config dict. Qwen3.5 stores this inside ``rope_parameters`` rather than at top level. """ prf = llm_dict.get("partial_rotary_factor") if prf is not None: return float(prf) for key in ("rope_parameters", "rope_scaling"): nested = llm_dict.get(key) if not isinstance(nested, dict): continue if nested.get("partial_rotary_factor") is not None: return float(nested["partial_rotary_factor"]) for attention_type in ("full_attention", "sliding_attention"): attention_params = nested.get(attention_type) if (isinstance(attention_params, dict) and attention_params.get("partial_rotary_factor") is not None): return float(attention_params["partial_rotary_factor"]) return 1.0 def _detect_has_qk_norm(model_dir: str, model_type: str = "") -> bool: """Detect QK-norm by scanning checkpoint key names. ``*.q_norm.weight`` (Qwen3 convention) is model-agnostic. The HunYuan V1 ``*.query_layernorm.weight`` spelling is only honored for model families whose loader remaps it onto ``q_norm`` (see ``_hunyuan_key_remap``); other families using that key name for unrelated norms must not trip the fused QK-norm path. """ keys = _checkpoint_weight_keys(model_dir) if any(".q_norm.weight" in k for k in keys): return True return (model_type in _QUERY_LAYERNORM_KEY_MODEL_TYPES and any(".query_layernorm.weight" in k for k in keys)) def _checkpoint_weight_keys(model_dir: str) -> List[str]: """Return checkpoint tensor keys, ignoring stale shard indexes when needed.""" index_path = os.path.join(model_dir, "model.safetensors.index.json") single_path = os.path.join(model_dir, "model.safetensors") if os.path.exists(index_path): try: with open(index_path) as f: index = json.load(f) weight_map = index.get("weight_map", {}) missing_shards = { shard for shard in set(weight_map.values()) if not os.path.exists(os.path.join(model_dir, shard)) } if not missing_shards or not os.path.exists(single_path): return list(weight_map.keys()) except (OSError, json.JSONDecodeError): pass if os.path.exists(single_path): try: from safetensors import safe_open with safe_open(single_path, framework="pt") as f: return list(f.keys()) except (OSError, ImportError): pass return [] _VL_LLM_PREFIXES = ("language_model.", "text_model.", "llm.", "thinker.") def _strip_vl_prefix(name: str) -> str: """Strip a known VL wrapper prefix (e.g. ``language_model.``) if present.""" if name.startswith("model.language_model."): return "model." + name[len("model.language_model."):] if name.startswith("model.visual."): return "visual." + name[len("model.visual."):] for prefix in _VL_LLM_PREFIXES: if name.startswith(prefix): return name[len(prefix):] return name def _normalize_module_name(name: str) -> str: """Normalise a checkpoint / hf_quant_config module name to the short namespace that ``make_linear`` uses (``layers.N...``, ``lm_head``, etc.). LLM ``make_linear`` callers pass names without any ``model.`` prefix (see ``modeling_default.py``: ``module_name=f"layers.{i}.mlp.gate_proj"``). For multimodal checkpoints whose keys are ``model.language_model.layers.N...`` the entire compound prefix must be stripped so the resulting short name matches what ``module_quant_type`` looks up. Single VL prefixes and bare ``model.`` follow the same rule. Used by ``_effective_excluded_modules``, ``_detect_modelopt_unquantized_linears``, and ``_parse_mixed_precision`` so that ``excluded``, ``layer_overrides``, and ``module_name`` all share the same name space. """ if name.startswith("model.language_model."): return name[len("model.language_model."):] if name.startswith("thinker.model."): return name[len("thinker.model."):] # Talker sub-LLM keys share the thinker short-name space after # ``_make_sub_model_dir(key_prefix='talker.')`` staging. if name.startswith("talker.model."): return name[len("talker.model."):] # Bare submodel prefixes (consolidated Qwen3-Omni roots): the visual / # audio / code_predictor builders use short names without them. if name.startswith("thinker."): return name[len("thinker."):] if name.startswith("talker."): return name[len("talker."):] if name.startswith("model.decoder."): return name[len("model.decoder."):] for prefix in _VL_LLM_PREFIXES + ("model.", ): if name.startswith(prefix): return name[len(prefix):] return name def _detect_unquantized_modules(model_dir: str) -> List[str]: """Return module names whose weights are plain float (not int4 quantized). Some layers (often ``lm_head``) use ``*.weight`` instead of ``*.qweight``. Checkpoint wrapper prefixes (``model.``, ``language_model.``, etc.) are stripped so that the returned names match the short names used by ``make_linear()``. """ all_keys = _checkpoint_weight_keys(model_dir) # Top-level module prefix = everything before the last dot segment qweight_modules = { k.rsplit(".", 1)[0] for k in all_keys if k.endswith(".qweight") } weight_modules = { k.rsplit(".", 1)[0] for k in all_keys if k.endswith(".weight") } excluded: List[str] = [ _normalize_module_name(m) for m in (weight_modules - qweight_modules) ] # lm_head may have neither .weight nor .qweight when tie_word_embeddings=True # (the checkpoint omits lm_head.weight entirely). Treat it as FP16 so that # tie_weights() can clone embed_tokens.weight into it after loading. all_linear_stripped = { _normalize_module_name(m) for m in qweight_modules | weight_modules } if "lm_head" not in all_linear_stripped: excluded.append("lm_head") return _with_gdn_fused_exclusions(excluded) _GDN_INPUT_PROJ_MODULES = ("in_proj_qkv", "in_proj_z", "in_proj_b", "in_proj_a") # Fused HF GDN projection names (as they appear in a ModelOpt ``ignore`` / # ``modules_to_not_convert`` list) mapped to the split projections # trt-edge-llm actually builds. The checkpoint loader splits # ``in_proj_qkvz`` -> ``in_proj_qkv`` / ``in_proj_z`` and ``in_proj_ba`` -> # ``in_proj_b`` / ``in_proj_a`` (Qwen3-Next family). Excluding only the fused # name would otherwise leave the split Linears at the dominant quant type # (e.g. NVFP4) even though their weights are plain FP16, so ``make_linear`` # would build an NVFP4Linear against an unquantized weight. _GDN_FUSED_PROJ_SPLITS: Dict[str, Tuple[str, ...]] = { "in_proj_qkvz": ("in_proj_qkv", "in_proj_z"), "in_proj_ba": ("in_proj_b", "in_proj_a"), } def _with_gdn_fused_exclusions(modules: List[str]) -> List[str]: """Reconcile GDN input-projection exclusions with the split/fused forms. Expands fused HF names (``in_proj_qkvz`` / ``in_proj_ba``) present in the ignore list into the split projections the model builds, and adds the synthetic ``in_proj_fused`` when all four split projections are FP16. """ result = set(modules) # Fused HF name -> split projections (so an excluded ``in_proj_qkvz`` # also excludes the ``in_proj_qkv`` / ``in_proj_z`` the model builds). for module in list(result): for fused, splits in _GDN_FUSED_PROJ_SPLITS.items(): suffix = f".{fused}" if module.endswith(suffix): prefix = module[:-len(suffix)] for split in splits: result.add(f"{prefix}.{split}") break by_prefix: Dict[str, set] = {} for module in result: for proj in _GDN_INPUT_PROJ_MODULES: suffix = f".{proj}" if module.endswith(suffix): by_prefix.setdefault(module[:-len(suffix)], set()).add(proj) break for prefix, projs in by_prefix.items(): if all(proj in projs for proj in _GDN_INPUT_PROJ_MODULES): result.add(f"{prefix}.in_proj_fused") return sorted(result) def _detect_quantized_modules(model_dir: str) -> List[str]: """Return modules that have checkpoint quantization sidecars.""" all_keys = _checkpoint_weight_keys(model_dir) suffixes = (".qweight", ".weight_scale", ".weight_scale_2", ".input_scale", ".scales") modules = { _strip_vl_prefix(k.rsplit(".", 1)[0]) for k in all_keys if k.endswith(suffixes) } return sorted(modules) def _effective_excluded_modules(model_dir: str, excluded: List[str]) -> List[str]: """Drop exclusions contradicted by quantized tensors in the checkpoint, and normalize remaining names so they match what ``make_linear`` looks up. Normalisation mirrors ``_parse_mixed_precision``'s strip on ``layer_overrides`` keys (``language_model.`` / ``text_model.`` / ``llm.`` / ``model.``) so a checkpoint that writes ``model.visual.blocks.X.Y`` to ``exclude_modules`` matches the ``visual.blocks.X.Y`` ``module_name`` the modeling code passes. """ quantized_modules = set(_detect_quantized_modules(model_dir)) normalized = [ _normalize_module_name(module) for module in excluded if _strip_vl_prefix(module) not in quantized_modules ] return _with_gdn_fused_exclusions(normalized) def _detect_modelopt_unquantized_linears(model_dir: str, submodel_prefix: str = "" ) -> List[str]: """Return module_name strings of Linears the ModelOpt checkpoint left unquantized. A ModelOpt-quantized Linear stores both ``<name>.weight`` (packed) and ``<name>.weight_scale`` (scale tensor; FP8 / NVFP4 / MXFP8 / AWQ / INT8-SQ all emit this). Linears that ModelOpt skipped — typically because their wildcard had ``enable: False`` (visual / audio / lm_head) — only have ``<name>.weight``. We compute the set of "has .weight without .weight_scale" modules from the checkpoint index, then return them in the short form ``make_linear`` uses (leading ``model.`` and VL wrapper prefixes stripped). Norm and embedding names also fall into this set but are harmless: ``make_linear`` only consults ``excluded`` for paths that actually go through it (i.e. real Linears), so extra entries are inert. Used by ``_parse_quant`` to plug a long-standing gap: for dominant-quant checkpoints (``quant_algo: FP8 / NVFP4 / W4A16_AWQ / ...``) ModelOpt does NOT populate ``exclude_modules`` even when whole submodules (visual tower, audio encoder) were skipped during PTQ. Without this augmentation, ``make_linear``'s dominant fallback would build NVFP4Linear / FP8Linear against FP16 weights → shape mismatch on export. Complements ``_effective_excluded_modules`` (which drops false-positive excludes when the checkpoint contradicts them). This helper supplies the opposite direction: false-negative excludes when ``exclude_modules`` is silent on a submodule that was actually skipped during PTQ. """ all_keys = _checkpoint_weight_keys(model_dir) if not all_keys: return [] # Omni roots hold every sub-model's tensors. Restrict the scan, or another # sub-model's unquantized Linears would be normalised onto this one's # module names and silently unquantize them. if submodel_prefix: all_keys = [k for k in all_keys if k.startswith(submodel_prefix)] weight_modules = { k.rsplit(".", 1)[0] for k in all_keys if k.endswith(".weight") } scale_modules = { k.rsplit(".", 1)[0] for k in all_keys if k.endswith(".weight_scale") } unquantized = weight_modules - scale_modules excluded = set() for name in unquantized: normalized = _normalize_module_name(name) if normalized.endswith(".self_attn.qkv_proj"): prefix = normalized[:-len("qkv_proj")] excluded.update(f"{prefix}{proj}" for proj in ("q_proj", "k_proj", "v_proj")) elif normalized.endswith(".mlp.gate_up_proj"): prefix = normalized[:-len("gate_up_proj")] excluded.update(f"{prefix}{proj}" for proj in ("gate_proj", "up_proj")) else: excluded.add(normalized) # Normalise to the same name space ``layer_overrides`` and ``excluded`` # use, so ``make_linear`` finds entries via its ``module_name`` lookup. return _with_gdn_fused_exclusions(sorted(excluded)) def _scope_exclusions(patterns: List[str], submodel_prefix: str) -> List[str]: """Keep only exclusions that apply to *submodel_prefix*, prefix stripped. Omni roots share one ignore list across thinker/talker, and short-name normalisation later drops both prefixes. A talker-scoped glob such as ``talker.model.layers.0.mlp.shared_expert*`` would therefore also unquantize the thinker's own shared_expert. A leading ``*`` marks a submodel-agnostic glob and is kept as-is. """ if not submodel_prefix: return patterns scoped = [] for pat in patterns: if pat.startswith(submodel_prefix): scoped.append(pat[len(submodel_prefix):]) elif pat.startswith("*") or pat == "": scoped.append(pat) return scoped def _parse_quant(model_dir: str, config: dict, submodel_prefix: str = "") -> QuantConfig: """Determine quantisation config from hf_quant_config.json or config.json. Checkpoint-provided W4A16 ``lm_head`` (``W4A16_NVFP4`` in ``quantized_layers``, e.g. Qwen3.6-35B-A3B-NVFP4) maps to :data:`QUANT_NVFP4_A16`. Excluded FP16/BF16 heads stay FP16; export does not invent packed NVFP4 for them. """ # ---- Sidecar hf_quant_config.json --------------------------------------- hf_path = os.path.join(model_dir, "hf_quant_config.json") if os.path.exists(hf_path): with open(hf_path) as f: hq = json.load(f) q = hq.get("quantization", {}) algo = (q.get("quant_algo") or "").upper() if "AWQ" in algo and "W4A16" in algo: # Drop exclusions that the checkpoint contradicts (false positives), # then augment with submodules ModelOpt actually left unquantized # (false negatives — typically visual tower / audio encoder). excluded = _effective_excluded_modules( model_dir, _scope_exclusions(list(q.get("exclude_modules", [])), submodel_prefix)) excluded.extend(m for m in _detect_modelopt_unquantized_linears( model_dir, submodel_prefix) if m not in excluded) return QuantConfig( quant_type=QUANT_INT4_AWQ_MODELOPT, group_size=int(q.get("group_size", 128)), excluded=excluded, ) if algo == "MIXED_PRECISION": quantized_layers = q.get("quantized_layers", {}) dominant, group_size, layer_overrides, layer_group_sizes = _parse_mixed_precision( quantized_layers, config.get("model_type") or "") return QuantConfig( quant_type=dominant, group_size=group_size, kv_cache_quant=_detect_llm_kv_cache_fp8(model_dir), visual_mha_quant=_detect_visual_mha_fp8(model_dir), excluded=_effective_excluded_modules( model_dir, _scope_exclusions(list(q.get("exclude_modules", [])), submodel_prefix)), layer_overrides=layer_overrides, layer_group_sizes=layer_group_sizes, is_mixed_precision=True, ) qt = _algo_to_quant_type(algo) gs = int(q.get("group_size", 1)) if qt == QUANT_MXFP8 and gs == 1: gs = 32 # MXFP8 default block_size excluded = _effective_excluded_modules( model_dir, _scope_exclusions(list(q.get("exclude_modules", [])), submodel_prefix)) excluded.extend(m for m in _detect_modelopt_unquantized_linears( model_dir, submodel_prefix) if m not in excluded) return QuantConfig( quant_type=qt, group_size=gs, kv_cache_quant=_detect_llm_kv_cache_fp8(model_dir), visual_mha_quant=_detect_visual_mha_fp8(model_dir), excluded=excluded, ) # ---- Embedded quantization_config in config.json ------------------------ qc = config.get("quantization_config") if qc is None: return QuantConfig() # Embedded block with ``quant_algo`` (export tool formats) if "quant_algo" in qc: algo = (qc.get("quant_algo") or "").upper() if algo == "MIXED_PRECISION": dominant, group_size, layer_overrides, layer_group_sizes = _parse_mixed_precision( qc.get("quantized_layers", {}), config.get("model_type") or "") return QuantConfig( quant_type=dominant, group_size=group_size, kv_cache_quant=_detect_llm_kv_cache_fp8(model_dir), visual_mha_quant=_detect_visual_mha_fp8(model_dir), excluded=_effective_excluded_modules( model_dir, _scope_exclusions( list(qc.get("exclude_modules", qc.get("ignore", []))), submodel_prefix)), layer_overrides=layer_overrides, layer_group_sizes=layer_group_sizes, is_mixed_precision=True, ) if "W4A16" in algo and "AWQ" in algo: return QuantConfig( quant_type=QUANT_INT4_AWQ_MODELOPT, group_size=int(qc.get("group_size", 128)), excluded=_effective_excluded_modules( model_dir, _scope_exclusions(list(qc.get("ignore", [])), submodel_prefix)), ) group_size = 1 cg = qc.get("config_groups", {}) if cg: first_group = next(iter(cg.values()), {}) group_size = int( first_group.get("weights", {}).get("group_size", 1)) elif qc.get("group_size") is not None: group_size = int(qc.get("group_size")) quant_type = _algo_to_quant_type(algo) # NVFP4 uses a fixed FP8-block group size of 16. Minimal ModelOpt # ``quantization_config`` blocks (``quant_algo`` only, no # ``config_groups`` / ``group_size`` — e.g. Qwen3-Omni Next NVFP4) # omit it, so default it here rather than leaving the per-tensor 1. if quant_type == QUANT_NVFP4 and group_size == 1: group_size = 16 kv = qc.get("kv_cache_scheme") kv_str = "fp8" if kv else None return QuantConfig( quant_type=quant_type, group_size=group_size, kv_cache_quant=kv_str, excluded=_effective_excluded_modules( model_dir, _scope_exclusions(list(qc.get("ignore", [])), submodel_prefix)), ) # quant_method == awq (column-packed int4 checkpoints) if qc.get("quant_method") == "awq": return QuantConfig( quant_type=QUANT_INT4_AWQ, group_size=int(qc.get("group_size", 128)), excluded=_detect_unquantized_modules(model_dir), ) # quant_method == gptq if qc.get("quant_method") == "gptq": return QuantConfig( quant_type=QUANT_INT4_GPTQ, group_size=int(qc.get("group_size", 128)), gptq_zero_point_offset=_detect_gptq_zero_point_offset( model_dir, qc), excluded=_detect_unquantized_modules(model_dir), ) return QuantConfig() def _detect_gptq_zero_point_offset(model_dir: str, qc: dict) -> int: """Infer whether GPTQ qzeros stores ``zero`` or ``zero - 1``. Older symmetric GPTQ checkpoints used by Qwen3 store packed ``0x77777777`` for a real zero point of 8. Qwen3.5 stores packed ``0x88888888`` for the same real zero point. Inspecting a tiny qzeros sample lets both variants share the same repacking path without depending on GPTQModel/Optimum. """ if not bool(qc.get("sym", False)): return 1 try: from safetensors import safe_open index_path = os.path.join(model_dir, "model.safetensors.index.json") if os.path.exists(index_path): with open(index_path) as f: weight_map: dict = json.load(f).get("weight_map", {}) qzeros_key = next((k for k in weight_map if k.endswith(".qzeros")), None) if qzeros_key is None: return 1 shard_path = os.path.join(model_dir, weight_map[qzeros_key]) else: shard_path = os.path.join(model_dir, "model.safetensors") if not os.path.exists(shard_path): return 1 with safe_open(shard_path, framework="pt") as f: qzeros_key = next( (k for k in f.keys() if k.endswith(".qzeros")), None) if qzeros_key is None: return 1 with safe_open(shard_path, framework="pt") as f: qzeros = f.get_tensor(qzeros_key).flatten()[:1024].cpu().tolist() except Exception: return 1 if not qzeros: return 1 nibbles = [] for value in qzeros: packed = int(value) & 0xFFFFFFFF nibbles.extend((packed >> (4 * i)) & 0xF for i in range(8)) if nibbles and all(v == 8 for v in nibbles): return 0 return 1 def _algo_to_quant_type(algo: str) -> str: algo = algo.upper() # MXFP8 per-block must be checked before generic FP8 if "FP8_PB" in algo or "MXFP8" in algo: return QUANT_MXFP8 if "FP8" in algo: return QUANT_FP8 if "FP4" in algo or "NVFP4" in algo: return QUANT_NVFP4 # W4A16_AWQ from ModelOpt unified checkpoints uses prepacked uint8 weights; # plain AWQ / INT4_AWQ from HuggingFace uses column-packed int32 qweight. if "W4A16" in algo and "AWQ" in algo: return QUANT_INT4_AWQ_MODELOPT if "AWQ" in algo or "INT4_AWQ" in algo: return QUANT_INT4_AWQ if "W8A8" in algo or "INT8" in algo: return QUANT_INT8_SQ return QUANT_FP16 def _parse_mixed_precision( quantized_layers: dict, model_type: str = "") -> "tuple[str, int, dict, dict]": """Parse MIXED_PRECISION quantized_layers dict. Returns ``(dominant_quant_type, dominant_group_size, layer_overrides, layer_group_sizes)``. ``layer_overrides`` maps **every** quantized module name to its quant-type string. Modules not listed in ``quantized_layers`` are unquantized (FP16); ``make_linear`` falls back to FP16 when a module_name is absent from ``layer_overrides``. """ from collections import Counter def _mixed_quant_type(algo: str) -> str: qt = _algo_to_quant_type(algo) if (qt == QUANT_NVFP4 and "W4A16" in algo.upper() and model_type.lower().startswith("nemotron_h")): return QUANT_NVFP4_A16 return qt def _group_size(layer_config: dict, quant_type: str) -> int: configured = int(layer_config.get("group_size", 1)) if configured != 1: return configured if quant_type == QUANT_MXFP8: return 32 if quant_type in (QUANT_NVFP4, QUANT_NVFP4_A16): return 16 return configured algo_count: Counter = Counter() algo_group_size: dict = {} for layer_cfg in quantized_layers.values(): algo = layer_cfg.get("quant_algo", "").upper() algo_count[algo] += 1 if algo not in algo_group_size: algo_group_size[algo] = _group_size(layer_cfg, _mixed_quant_type(algo)) if not algo_count: return QUANT_FP16, 1, {}, {} dominant_algo = algo_count.most_common(1)[0][0] dominant_type = _mixed_quant_type(dominant_algo) dominant_group_size = algo_group_size.get(dominant_algo, 1) # Expand fused projection keys (``self_attn.qkv_proj``, # ``mlp.gate_up_proj``) into the split names ``make_linear`` looks up # (``q_proj``/``k_proj``/``v_proj`` and ``gate_proj``/``up_proj``). layer_overrides: dict = {} layer_group_sizes: dict = {} for name, layer_cfg in quantized_layers.items(): algo = layer_cfg.get("quant_algo", "").upper() short_name = _normalize_module_name(name) quant_type = _mixed_quant_type(algo) group_size = _group_size(layer_cfg, quant_type) if (_is_muse_glimmer_model_type(model_type) and short_name.endswith(".self_attn.output_gate_proj")): short_name = short_name[:-len(".output_gate_proj")] + ".gate_proj" if short_name.endswith(".self_attn.qkv_proj"): prefix = short_name[:-len("qkv_proj")] module_names = tuple(f"{prefix}{proj}" for proj in ("q_proj", "k_proj", "v_proj")) elif short_name.endswith(".mlp.gate_up_proj"): prefix = short_name[:-len("gate_up_proj")] module_names = tuple(f"{prefix}{proj}" for proj in ("gate_proj", "up_proj")) else: module_names = (short_name, ) for module_name in module_names: layer_overrides[module_name] = quant_type layer_group_sizes[module_name] = group_size return (dominant_type, dominant_group_size, layer_overrides, layer_group_sizes) def _kv_norm(s: Optional[str]) -> Optional[str]: if not s: return None return s.strip().lower() or None # Keep in sync with quantization_configs._VISUAL_PREFIXES (not imported here: # that module pulls in modelopt, too heavy for config parsing). _VISUAL_PATH_HINTS = ("visual", "vision_tower", "vision_model", "multi_modal_projector", "mlp1", "image_embed", "embed_vision") def _safetensor_keys(model_dir: str) -> Optional[List[str]]: """Return all tensor keys from sharded or single-file safetensors.""" index_path = os.path.join(model_dir, "model.safetensors.index.json") if os.path.exists(index_path): with open(index_path) as f: return list(json.load(f).get("weight_map", {}).keys()) st_path = os.path.join(model_dir, "model.safetensors") if not os.path.exists(st_path): return None from safetensors import safe_open with safe_open(st_path, framework="pt") as f: return list(f.keys()) def _detect_visual_mha_fp8(model_dir: str) -> Optional[str]: """Return ``"fp8"`` iff the checkpoint carries visual-MHA Q/K/V scales. Self-describing detection: the presence of ``<visual_prefix>*.q_scale`` buffers in the safetensors is the ground-truth signal that ViT FP8 MHA was calibrated. Avoids trusting modelopt's hf_quant_config.json metadata (which has no ``visual_mha_quant_algo`` slot and conflates ViT q/k/v_bmm with the LLM KV cache). """ keys = _safetensor_keys(model_dir) if keys is None: return None for k in keys: if k.endswith(".q_scale") and any(h in k for h in _VISUAL_PATH_HINTS): return "fp8" return None def _detect_llm_kv_cache_fp8(model_dir: str) -> Optional[str]: """Return ``"fp8"`` iff the checkpoint carries LLM KV-cache K/V scales. Self-describing detection: an ``<llm_attn>.k_proj.k_scale`` buffer outside of any visual / vision subtree is the ground-truth signal. Modelopt's ``hf_quant_config.json:kv_cache_quant_algo`` field would say "FP8" even when only ViT MHA was requested (because its writer trips on any enabled ``k_bmm_quantizer``, including the visual one); ignoring that field and looking at the checkpoint directly avoids the false positive. """ keys = _safetensor_keys(model_dir) if keys is None: return None for k in keys: if k.endswith(".k_proj.k_scale") and not any( h in k for h in _VISUAL_PATH_HINTS): return "fp8" return None