import glob import importlib import importlib.util import inspect import json import logging import math import struct import warnings from dataclasses import dataclass, fields from io import BytesIO from pathlib import Path from textwrap import dedent from typing import Any, Dict, List, Optional, Tuple, Union import mlx.core as mx import mlx.nn as nn import numpy as np import requests from huggingface_hub import snapshot_download from mlx.utils import tree_flatten, tree_map from PIL import Image, ImageOps from transformers import AutoProcessor from transformers.processing_utils import ProcessorMixin from .models.base import BaseImageProcessor from .quantization.one_bit import _quantization_for_path, replace_one_bit_modules from .tokenizer_utils import load_tokenizer from .trainer.utils import apply_lora_layers logger = logging.getLogger(__name__) # Modes that support activation quantization ACTIVATION_QUANTIZATION_MODES = {"nvfp4", "mxfp8"} # Constants MODEL_REMAPPING = { "moondream1": "moondream2", "llava_qwen2": "fastvlm", # Apple's FastVLM, note it's different to the one below "llava-qwen2": "llava_bunny", "bunny-llama": "llava_bunny", "lfm2-vl": "lfm2_vl", "cohere2_vision": "aya_vision", "jvlm": "jina_vlm", "phi4-siglip": "phi4_siglip", "sam3_video": "sam3", "sam3.1_video": "sam3_1", "granite-vision": "granite_vision", "granite4-vision": "granite4_vision", "granite4_vision": "granite4_vision", "rf-detr": "rfdetr", "dinov2_with_registers": "dinov2", "falcon-perception": "falcon_perception", "nemotronh_nano_omni_reasoning_v3": "nemotron_h_nano_omni", "cohere2moe": "cohere2_moe", "unlimited-ocr": "unlimited_ocr", "mistral": "llama", "phi-msft": "phixtral", "falcon_mamba": "mamba", "joyai_llm_flash": "deepseek_v3", "kimi_k2": "deepseek_v3", "minimax_m2": "minimax", "iquestcoder": "llama", "nemotron-nas": "nemotron_nas", "inkling_mm_model": "inkling", "lille-130m": "lille_130m", } MAX_FILE_SIZE_GB = 5 MODEL_CONVERSION_DTYPES = ["float16", "bfloat16", "float32"] SAFETENSORS_DTYPE_FALLBACKS = {"F8_E8M0": "U8"} GENERATION_CONFIG_DEFAULT_KEYS = ( "eos_token_id", "temperature", "top_p", "top_k", "do_sample", ) def apply_generation_config_defaults(model_config, config: dict): for key in GENERATION_CONFIG_DEFAULT_KEYS: if key in config: setattr(model_config, key, config[key]) return model_config def _merge_generation_config(config: dict, generation_config: dict) -> None: if not isinstance(generation_config, dict) or not generation_config: return config["generation_config"] = generation_config for key in GENERATION_CONFIG_DEFAULT_KEYS: if key in generation_config: config[key] = generation_config[key] def _e4m3_decode_table() -> mx.array: """Return a 256-entry ``float32`` LUT mapping every E4M3FN byte to its value. OCP ``E4M3FN``: 1 sign / 4 exponent (bias 7) / 3 mantissa bits, no infinities, and ``(exp=15, mant=7)`` reserved for NaN (max finite magnitude is therefore 448). This matches the byte convention MLX uses for ``nvfp4`` scales (verified empirically). """ table = [] for byte in range(256): sign = (byte >> 7) & 1 exponent = (byte >> 3) & 0xF mantissa = byte & 0x7 if exponent == 0: value = (mantissa / 8.0) * 2.0**-6 # subnormal elif exponent == 15 and mantissa == 7: value = float("nan") else: value = (1.0 + mantissa / 8.0) * 2.0 ** (exponent - 7) table.append(-value if sign else value) return mx.array(table, dtype=mx.float32) # Built once; reused by every NVFP4 fold. _E4M3_DECODE_LUT = _e4m3_decode_table() def _f32_to_e4m3(x: mx.array) -> mx.array: """Encode non-negative ``float32`` values to ``E4M3FN`` bytes. Pure-MLX bit manipulation (MLX exposes no float8 dtype). Saturates to 448 on overflow and flushes to the subnormal grid / zero on underflow. Inputs are assumed ``>= 0`` (NVFP4 group scales are magnitudes), so the sign bit is always 0. """ x = mx.maximum(x.astype(mx.float32), 0.0) bits = x.view(mx.uint32) fexp = (bits >> 23) & 0xFF # fp32 exponent, bias 127 fman = bits & 0x7FFFFF # fp32 mantissa, 23 bits # Normal path: target E4M3 biased exponent e = (fexp - 127) + 7. exponent = fexp.astype(mx.int32) - 120 drop = 20 # 23 -> 3 mantissa bits round_bit = (fman >> (drop - 1)) & 1 sticky = (fman & ((1 << (drop - 1)) - 1)) != 0 mantissa = fman >> drop roundup = round_bit & (sticky.astype(mx.uint32) | (mantissa & 1)) mantissa = mantissa + roundup carry = mantissa >> 3 # mantissa overflowed past 7 -> bump exponent mantissa = mantissa & 0x7 exponent = exponent + carry.astype(mx.int32) # Saturate: e > 15, or the NaN slot (e == 15, mant == 7), clamps to 448. over = (exponent > 15) | ((exponent == 15) & (mantissa == 7)) exponent = mx.where(over, mx.array(15, mx.int32), exponent) mantissa = mx.where(over, mx.array(6, mx.uint32), mantissa) normal_byte = (exponent.astype(mx.uint32) << 3) | mantissa normal_valid = exponent >= 1 # Subnormal path: value = m * 2^-9, so m = round(x * 512) (RNE). # m == 8 lands exactly on the smallest normal (0x08 = e1 m0 = 2^-6). sub = x * 512.0 sub_floor = mx.floor(sub) frac = sub - sub_floor sub_floor_u32 = sub_floor.astype(mx.uint32) up = (frac > 0.5) | ((frac == 0.5) & ((sub_floor_u32 & 1) == 1)) sub_byte = sub_floor_u32 + up.astype(mx.uint32) byte = mx.where(normal_valid, normal_byte, sub_byte) return byte.astype(mx.uint8) def _transform_modelopt_nvfp4_weights( weights: Dict[str, mx.array], quantization_config: Optional[Dict[str, Any]], ) -> Tuple[Dict[str, mx.array], Optional[Dict[str, Any]]]: """Convert ModelOpt NVFP4 and mixed NVFP4/FP8 checkpoints. ModelOpt's mixed export uses ``weight_scale_2`` to identify NVFP4 linears and a lone ``weight_scale`` for FP8 linears. MLX can load the former natively; the latter are decoded to dense weights because its FP8 mode uses block scales rather than ModelOpt's per-tensor/channel scale. """ if quantization_config is None: return weights, None if quantization_config.get("quant_method") not in { "modelopt", "modelopt_mixed", } or quantization_config.get("quant_algo") not in { "NVFP4", "W4A16_NVFP4", "MIXED_PRECISION", }: return weights, None scale_2_suffix = ".weight_scale_2" nvfp4_prefixes = { key[: -len(scale_2_suffix)] for key in weights if key.endswith(scale_2_suffix) } scale_suffix = ".weight_scale" scaled_prefixes = { key[: -len(scale_suffix)] for key in weights if key.endswith(scale_suffix) } fp8_prefixes = scaled_prefixes - nvfp4_prefixes if not nvfp4_prefixes and not fp8_prefixes: return weights, None nvfp4_consumed = { f"{prefix}.{suffix}" for prefix in nvfp4_prefixes for suffix in ("weight", "weight_scale", "input_scale") } fp8_consumed = { f"{prefix}.{suffix}" for prefix in fp8_prefixes for suffix in ("weight", "input_scale") } transformed = {} # Each fold below builds a deep lazy graph. A large MoE export has tens of # thousands of quantized tensors, so the unevaluated intermediates blow past # Metal's live-buffer limit before the dict is ever consumed. Flush in # batches to keep the graph shallow; this also frees the intermediates. pending: List[mx.array] = [] def _flush(force: bool = False) -> None: if pending and (force or len(pending) >= 256): mx.eval(pending) pending.clear() for key, value in weights.items(): # ModelOpt emits per-layer FP8 KV-cache scales when kv_cache_quant_algo # is set. MLX quantizes the KV cache at runtime and has no parameter to # hold them, so drop them rather than fail the strict load. if key.endswith(".k_scale") or key.endswith(".v_scale"): continue if key.endswith(scale_2_suffix): prefix = key[: -len(scale_2_suffix)] weight_key = f"{prefix}.weight" scale_key = f"{prefix}.weight_scale" if weight_key not in weights or scale_key not in weights: raise ValueError(f"Missing ModelOpt NVFP4 tensors for {prefix}.") weight = weights[weight_key] scale = weights[scale_key] if ( weight.dtype != mx.uint8 or scale.dtype != mx.uint8 or weight.ndim != 2 or scale.ndim != 2 or value.size != 1 ): raise ValueError(f"Invalid ModelOpt NVFP4 tensors for {prefix}.") if ( weight.shape[0] != scale.shape[0] or weight.shape[1] != 8 * scale.shape[1] ): raise ValueError(f"Invalid ModelOpt NVFP4 scale shape for {prefix}.") transformed[weight_key] = weight.view(mx.uint32) decoded_scale = _E4M3_DECODE_LUT[scale.astype(mx.uint32)] transformed[f"{prefix}.scales"] = _f32_to_e4m3( decoded_scale * value.astype(mx.float32) ) pending.append(transformed[f"{prefix}.scales"]) _flush() elif key.endswith(scale_suffix) and key[: -len(scale_suffix)] in fp8_prefixes: prefix = key[: -len(scale_suffix)] weight_key = f"{prefix}.weight" if weight_key not in weights or weights[weight_key].dtype != mx.uint8: raise ValueError(f"Invalid ModelOpt FP8 tensors for {prefix}.") if not mx.issubdtype(value.dtype, mx.floating): raise ValueError(f"Invalid ModelOpt FP8 scale for {prefix}.") transformed[weight_key] = _dequantize_compressed_tensors_fp8_weight( weights[weight_key], value ) pending.append(transformed[weight_key]) _flush() elif key in nvfp4_consumed or key in fp8_consumed: continue else: transformed[key] = value _flush(force=True) quantization = None if nvfp4_prefixes: quantization = {"group_size": 16, "bits": 4, "mode": "nvfp4"} return transformed, quantization def _transform_compressed_tensors_nvfp4_weights( weights: Dict[str, mx.array], quantization_config: Dict[str, Any], ) -> Dict[str, mx.array]: """Fold compressed-tensors NVFP4 weights into MLX-native ``nvfp4`` weights. A ``nvfp4-pack-quantized`` checkpoint stores, per quantized Linear: - ``
.weight_packed`` ``uint8`` ``[out, in // 2]`` (2x E2M1 per byte) - ``
.weight_scale`` ``uint8`` ``[out, in // 16]`` (E4M3 per group of 16, loaded by ``mx.load`` as raw bytes -- the same byte layout MLX uses for nvfp4 scales) - ``
.weight_global_scale`` ``float32`` ``[1]`` (per-tensor; the real weight is ``fp4 * weight_scale / weight_global_scale``) MLX ``nvfp4`` ``QuantizedLinear`` expects ``
.weight`` (``uint32``) plus ``
.scales`` (``uint8`` E4M3) and is single-level: the per-tensor global scale is not representable (and is rejected on the Metal backend). Both decodes are linear in the FP4 codes, so the global scale can be folded directly into the per-group E4M3 scales: ``scale_mlx = E4M3(weight_scale / global_scale)``. We keep the original packed E2M1 codes bit-exact, avoiding the weight dequantize/re-quantize round-trip entirely. """ packed_suffix = ".weight_packed" new_weights = {} for key in list(weights.keys()): if key.endswith(packed_suffix): prefix = key[: -len(packed_suffix)] packed = weights[key] scale = weights[f"{prefix}.weight_scale"] global_scale = weights[f"{prefix}.weight_global_scale"].astype(mx.float32) # weight_packed is uint8 [out, in//2]; reinterpret as uint32 # [out, in//8] to match MLX's nvfp4 layout (bit-identical). new_weights[f"{prefix}.weight"] = packed.view(mx.uint32) # Fold the per-tensor global scale into the per-group E4M3 scales: # decode E4M3 -> divide by global scale -> re-encode E4M3. # The FP4 codes are untouched; only the much smaller scale tensor # is re-rounded once. decoded = _E4M3_DECODE_LUT[scale.astype(mx.uint32)] new_weights[f"{prefix}.scales"] = _f32_to_e4m3(decoded / global_scale) elif key.endswith(".weight_scale") or key.endswith(".weight_global_scale"): # Consumed alongside their ``.weight_packed``. continue else: new_weights[key] = weights[key] return new_weights def _transform_compressed_tensors_int4_weights( weights: Dict[str, mx.array], quantization_config: Dict[str, Any], ) -> Tuple[Dict[str, mx.array], Dict[str, Any]]: """Remap compressed-tensors INT4 ``pack-quantized`` weights to MLX affine. A symmetric int4 ``pack-quantized`` checkpoint stores, per quantized Linear: - ``
.weight_packed`` ``int32`` ``[out, in // 8]`` (8x int4 per word, LSB-first) - ``
.weight_scale`` (``bf16``/``float``) ``[out, in // group_size]`` - ``
.weight_shape`` ``int64`` ``[2]`` (unused by MLX)
MLX affine ``QuantizedLinear`` uses the same int4 packing and dequantizes as
``w * scale + bias``. Symmetric int4 stores values in ``[0, 15]``
representing ``[-8, 7]``, i.e. ``value = packed - 8``, so we set
``bias = -8 * scale``. The packed ``int32`` is bit-identical to MLX's
``uint32`` layout (reinterpreted via ``view``).
"""
weights_cfg = (
quantization_config.get("config_groups", {})
.get("group_0", {})
.get("weights", {})
)
group_size = weights_cfg.get("group_size", 32)
bits = weights_cfg.get("num_bits", 4)
packed_suffix = ".weight_packed"
new_weights = {}
for key in list(weights.keys()):
if key.endswith(packed_suffix):
prefix = key[: -len(packed_suffix)]
scale = weights[f"{prefix}.weight_scale"]
new_weights[f"{prefix}.weight"] = weights[key].view(mx.uint32)
new_weights[f"{prefix}.scales"] = scale
new_weights[f"{prefix}.biases"] = -(2 ** (bits - 1)) * scale
elif key.endswith(".weight_scale") or key.endswith(".weight_shape"):
# Consumed alongside their ``.weight_packed`` (shape is unused).
continue
else:
new_weights[key] = weights[key]
return new_weights, {"group_size": group_size, "bits": bits, "mode": "affine"}
# Compressed-tensors auxiliary tensors with no MLX consumer. They carry
# activation / kv-cache / shape metadata; left in place they reach
# ``model.load_weights(strict=True)`` as unexpected keys and abort startup.
_COMPRESSED_TENSORS_DROP_SUFFIXES = (
".weight_shape",
".weight_global_scale", # folded into ``.scales`` alongside ``.weight_packed``
".weight_zero_point",
".input_global_scale",
".input_scale",
".k_scale",
".v_scale",
)
def _compressed_tensors_group_weights(
quantization_config: Dict[str, Any], ct_format: str
) -> Dict[str, Any]:
"""Return the ``weights`` sub-config of the first group with ``ct_format``."""
for group in (quantization_config.get("config_groups") or {}).values():
if isinstance(group, dict) and group.get("format") == ct_format:
return group.get("weights") or {}
return {}
def _dequantize_compressed_tensors_fp8_weight(
weight: mx.array,
scale: mx.array,
block_structure: Optional[List[int]] = None,
) -> mx.array:
"""Dequantize an fp8 ``float-quantized`` weight to dense.
``float-quantized`` stores the weight as ``float8_e4m3fn`` -- which
``mx.load`` surfaces as raw ``uint8`` E4M3 bytes -- plus either a
per-output-channel or blockwise ``weight_scale``. MLX has no matching FP8
mode, so decode the E4M3 codes and rescale into a dense tensor. The dense
weight is emitted in ``weight_scale``'s dtype (typically ``bfloat16``).
"""
if weight.dtype != mx.uint8 or weight.ndim < 2:
raise ValueError(
"Compressed-tensors FP8 weights must be E4M3 byte matrices; "
f"got dtype={weight.dtype}, shape={weight.shape}."
)
out_dtype = scale.dtype if scale.dtype != mx.float32 else mx.bfloat16
decoded = _E4M3_DECODE_LUT[weight.astype(mx.uint32)] # float32 [out, in]
scale = scale.astype(mx.float32)
if block_structure is None:
if scale.ndim == 1:
scale = scale[:, None]
return (decoded * scale).astype(out_dtype)
if len(block_structure) != 2 or any(size <= 0 for size in block_structure):
raise ValueError(
"Compressed-tensors FP8 block_structure must contain two positive "
f"dimensions; got {block_structure}."
)
block_rows, block_cols = block_structure
*batch_shape, rows, cols = weight.shape
row_blocks = (rows + block_rows - 1) // block_rows
col_blocks = (cols + block_cols - 1) // block_cols
expected_scale_shape = (*batch_shape, row_blocks, col_blocks)
if scale.shape != expected_scale_shape:
raise ValueError(
"Compressed-tensors FP8 scale shape does not match its weight: "
f"weight={weight.shape}, scales={scale.shape}, "
f"block_structure={block_structure}, expected={expected_scale_shape}."
)
pad_rows = row_blocks * block_rows - rows
pad_cols = col_blocks * block_cols - cols
if pad_rows or pad_cols:
decoded = mx.pad(
decoded,
[(0, 0)] * len(batch_shape) + [(0, pad_rows), (0, pad_cols)],
)
decoded = decoded.reshape(
*batch_shape,
row_blocks,
block_rows,
col_blocks,
block_cols,
)
decoded = (decoded * scale[..., :, None, :, None]).reshape(
*batch_shape,
rows + pad_rows,
cols + pad_cols,
)
return decoded[..., :rows, :cols].astype(out_dtype)
def _transform_compressed_tensors_mixed_weights(
weights: Dict[str, mx.array],
quantization_config: Dict[str, Any],
) -> Tuple[Dict[str, mx.array], Optional[Dict[str, Any]]]:
"""Route compressed-tensors mixed-precision or pure FP8 weights.
A ``mixed-precision`` export keeps a single top-level ``format`` and puts
the real formats in per-group ``config_groups``. A pure ``float-quantized``
export uses the same FP8 tensor layout without any packed weights. Rather
than match each group's regex ``targets``/``ignore`` against module paths
(fragile -- the HF names differ from MLX's), route every quantized Linear
by the tensors it actually carries, which is exactly what the group
assignment produced:
- ``.weight_packed`` + ``.weight_global_scale`` -> NVFP4
(folded to MLX-native ``nvfp4``, as ``_transform_..._nvfp4_weights`` does)
- ``.weight_packed`` alone -> INT4 ``pack-quantized`` (folded to ``affine``)
- ``.weight_scale`` without ``.weight_packed`` -> channel-wise fp8
``float-quantized`` (dequantized to a dense weight)
Layers under ``ignore`` (vision, ``linear_attn``, ``mtp``) carry a plain
``.weight`` with no scale and pass through untouched, so the later
``nn.quantize`` pass -- which only quantizes modules that have a ``.scales``
tensor -- leaves them dense.
Only one MLX-native quantized mode may coexist with the dense/fp8 layers
(``nn.quantize`` applies a single ``group_size``/``bits``/``mode``). The
reported models pair NVFP4 with dense fp8, so this holds; a checkpoint that
mixes e.g. NVFP4 and INT4 raises rather than silently mis-quantizing.
"""
packed_prefixes = {
key[: -len(".weight_packed")]
for key in weights
if key.endswith(".weight_packed")
}
# A ``.weight_scale`` with no packed sibling is a float-quantized fp8 layer.
fp8_prefixes = {
key[: -len(".weight_scale")] for key in weights if key.endswith(".weight_scale")
} - packed_prefixes
int4_cfg = _compressed_tensors_group_weights(quantization_config, "pack-quantized")
int4_bits = int4_cfg.get("num_bits", 4)
int4_group_size = int4_cfg.get("group_size", 32)
fp8_cfg = _compressed_tensors_group_weights(quantization_config, "float-quantized")
fp8_block_structure = fp8_cfg.get("block_structure")
new_weights: Dict[str, mx.array] = {}
native_quant: Dict[str, Dict[str, Any]] = {}
pending: List[mx.array] = []
def flush(force: bool = False) -> None:
if pending and (force or len(pending) >= 256):
mx.eval(pending)
pending.clear()
for key in list(weights):
if key not in weights:
continue
value = weights[key]
if key.endswith(".weight_packed"):
prefix = key[: -len(".weight_packed")]
scale = weights[f"{prefix}.weight_scale"]
global_key = f"{prefix}.weight_global_scale"
if global_key in weights: # NVFP4
global_scale = weights[global_key].astype(mx.float32)
new_weights[f"{prefix}.weight"] = value.view(mx.uint32)
decoded = _E4M3_DECODE_LUT[scale.astype(mx.uint32)]
new_weights[f"{prefix}.scales"] = _f32_to_e4m3(decoded / global_scale)
native_quant["nvfp4"] = {"group_size": 16, "bits": 4, "mode": "nvfp4"}
else: # INT4 symmetric pack-quantized
new_weights[f"{prefix}.weight"] = value.view(mx.uint32)
new_weights[f"{prefix}.scales"] = scale
new_weights[f"{prefix}.biases"] = -(2 ** (int4_bits - 1)) * scale
native_quant["affine"] = {
"group_size": int4_group_size,
"bits": int4_bits,
"mode": "affine",
}
continue
if key.endswith(".weight_scale"):
prefix = key[: -len(".weight_scale")]
if prefix in fp8_prefixes:
weight_key = f"{prefix}.weight"
source_weight = weights.pop(weight_key)
source_scale = weights.pop(key)
new_weights[f"{prefix}.weight"] = (
_dequantize_compressed_tensors_fp8_weight(
source_weight,
source_scale,
fp8_block_structure,
)
)
pending.append(new_weights[f"{prefix}.weight"])
flush()
# NVFP4 / INT4 scales are consumed with their ``.weight_packed``.
continue
if key.endswith(".weight") and key[: -len(".weight")] in fp8_prefixes:
continue # emitted (dequantized to dense) at its ``.weight_scale``
if any(key.endswith(suffix) for suffix in _COMPRESSED_TENSORS_DROP_SUFFIXES):
continue
new_weights[key] = value
flush(force=True)
if not native_quant:
return new_weights, None
if len(native_quant) > 1:
raise NotImplementedError(
"mixed-precision compressed-tensors checkpoint mixes multiple "
f"MLX-native quantized modes {sorted(native_quant)}; only one native "
"mode alongside dense/fp8 layers is supported."
)
return new_weights, next(iter(native_quant.values()))
def _transform_compressed_tensors_weights(
weights: Dict[str, mx.array],
quantization_config: Optional[Dict[str, Any]],
) -> Tuple[Dict[str, mx.array], Optional[Dict[str, Any]]]:
"""Transform compressed-tensors weights before model-specific sanitization.
This runs before MoE expert stacking so per-expert tensors have already been
renamed to MLX-native ``.weight`` / ``.scales`` / ``.biases`` keys.
"""
if quantization_config is None:
return weights, None
if quantization_config.get("quant_method") != "compressed-tensors":
return weights, None
# Mixed-precision and pure float-quantized exports both use
# ``.weight``/``.weight_scale`` for FP8 tensors. Route these before the
# packed-weight guard because a pure FP8 checkpoint has no
# ``.weight_packed`` tensors at all.
if quantization_config.get("format") in {
"mixed-precision",
"float-quantized",
}:
return _transform_compressed_tensors_mixed_weights(weights, quantization_config)
if not any(key.endswith(".weight_packed") for key in weights):
return weights, None
weights_cfg = (
quantization_config.get("config_groups", {})
.get("group_0", {})
.get("weights", {})
)
ct_format = quantization_config.get("format") or quantization_config.get(
"config_groups", {}
).get("group_0", {}).get("format")
quant_type = weights_cfg.get("type")
if ct_format == "nvfp4-pack-quantized":
return _transform_compressed_tensors_nvfp4_weights(
weights, quantization_config
), {"group_size": 16, "bits": 4, "mode": "nvfp4"}
if ct_format == "pack-quantized" and quant_type == "int":
return _transform_compressed_tensors_int4_weights(weights, quantization_config)
return weights, None
def quantize_activations(model: nn.Module) -> nn.Module:
def _maybe_qq(m: nn.Module) -> nn.Module:
"""Convert a QuantizedLinear layer to QQLinear if compatible."""
if isinstance(m, nn.QuantizedLinear):
if m.mode not in ACTIVATION_QUANTIZATION_MODES:
raise ValueError(
f"Mode ({m.mode}) does not support activation quantization. "
f"Supported modes: {', '.join(ACTIVATION_QUANTIZATION_MODES)}"
)
if m.get("bias", False):
raise ValueError(
"Linear layer with bias does not support activation quantization"
)
# Get dimensions from quantized weight
out_dims, in_dims = m.weight.shape
in_dims *= 32 // m.bits
qq = nn.QQLinear(in_dims, out_dims, m.group_size, m.bits, m.mode)
qq.quantize()
return qq
return m
leaves = tree_map(_maybe_qq, model.leaf_modules(), is_leaf=nn.Module.is_module)
model.update_modules(leaves)
return model
def skip_multimodal_module(path: str) -> bool:
"""
Check if a multimodal module (vision/audio) should skip quantization.
Args:
path: The module path to check
Returns:
bool: True if the module is multimodal and should skip quantization, False otherwise
"""
multimodal_modules = (
"vision",
"aligner",
"vision_model",
"vision_tower",
"vl_connector",
"sam_model",
"audio_model",
"audio_tower",
"code_predictor",
"img_projector",
"multi_modal_projector",
"patch_merge_mlp",
)
return any(module in path for module in multimodal_modules)
def _has_quantized_weights(path: str, weights: Optional[Dict[str, mx.array]]) -> bool:
return weights is not None and f"{path}.scales" in weights
def get_class_predicate(skip_vision=False, weights=None, quantization_config=None):
def predicate(p, m):
if (
skip_multimodal_module(p)
and skip_vision
and not _has_quantized_weights(p, weights)
):
return False
if quantization_config is not None and p in quantization_config:
return quantization_config[p]
if not hasattr(m, "to_quantized"):
return False
if hasattr(m, "weight") and m.weight.size % 64 != 0:
return False
if weights is not None:
return f"{p}.scales" in weights
return True
return predicate
def _language_model_quantization_config(model):
"""The model's own make_quantization_config, if its language module defines one."""
language_model = getattr(model, "language_model", None)
if language_model is None:
return None
module = importlib.import_module(type(language_model).__module__)
return getattr(module, "make_quantization_config", None)
def get_model_and_args(config: dict, model_path: Optional[Path] = None):
"""Resolve a model package and its normalized model type.
If the config declares ``model_file`` and ``model_path`` is provided, the
model module is imported from that file inside the checkpoint (the same
mechanism mlx_lm supports) instead of the built-in registry.
"""
if model_path is not None and (model_file := config.get("model_file")):
model_file_path = Path(model_path) / model_file
if not model_file_path.is_file():
raise FileNotFoundError(
f"config.json declares model_file={model_file!r} but "
f"{model_file_path} does not exist"
)
spec = importlib.util.spec_from_file_location("custom_model", model_file_path)
arch = importlib.util.module_from_spec(spec)
spec.loader.exec_module(arch)
return arch, "custom"
raw_model_type = config.get("model_type") or config.get("speculators_model_type")
if raw_model_type is None:
raise KeyError("model_type")
model_type = raw_model_type.lower()
model_type = MODEL_REMAPPING.get(model_type, model_type)
architectures = set(config.get("architectures") or ())
dflash_config = config.get("dflash_config")
if "Lfm2BidirectionalForMaskedLM" in architectures:
model_type = "lfm2_encoder"
elif "BoundaryExtractor" in architectures:
model_type = "gliner2_5"
elif "DFlash2DraftModel" in architectures:
model_type = "dflash2"
elif "Gemma4DSparkModel" in architectures:
model_type = "gemma4_dspark"
elif dflash_config is not None:
is_dspark = (
dflash_config.get("projector_type") == "dspark"
or int(config.get("markov_rank") or dflash_config.get("markov_rank") or 0)
> 0
)
model_type = "dspark" if is_dspark else f"{model_type}_dflash"
last_err: Optional[ImportError] = None
for pkg in ("mlx_vlm.models", "mlx_vlm.speculative.drafters"):
try:
arch = importlib.import_module(f"{pkg}.{model_type}")
return arch, model_type
except ImportError as e:
if model_type not in str(e):
raise
last_err = e
continue
msg = f"Model type {model_type} not supported. Error: {last_err}"
logging.error(msg)
raise ValueError(msg)
def _quantization_path_aliases(
path: str, model: Optional[nn.Module] = None
) -> Tuple[str, ...]:
"""Return checkpoint quantization keys that may refer to a module path."""
aliases = [path]
if path.startswith("language_model."):
aliases.append(path[len("language_model.") :])
model_aliases = getattr(model, "quantization_path_aliases", None)
if callable(model_aliases):
aliases.extend(model_aliases(path))
return tuple(dict.fromkeys(aliases))
def _quantization_for_module_path(
quantization: dict, path: str, model: Optional[nn.Module] = None
) -> Optional[dict]:
for alias in _quantization_path_aliases(path, model):
value = quantization.get(alias)
if isinstance(value, dict):
return value
if value is False:
return {}
return None
def _drop_modules_without_weights(
model: nn.Module, weights: dict, declared_keys: Optional[set] = None
) -> None:
"""Drop weightless top-level VLM modules the checkpoint manifest also omits."""
weighted_modules = {key.partition(".")[0] for key in weights}
declared_modules = {key.partition(".")[0] for key in (declared_keys or weights)}
dropped_modules = []
for name, child in list(model.items()):
if name == "language_model" or not isinstance(child, nn.Module):
continue
if not tree_flatten(child.parameters()) or name in weighted_modules:
continue
if name in declared_modules:
continue
setattr(model, name, None)
dropped_modules.append(name)
if dropped_modules:
logging.warning(
"Checkpoint has no weights for module(s): %s. Disabling them.",
", ".join(dropped_modules),
)
def get_model_path(
path_or_hf_repo: str,
revision: Optional[str] = None,
force_download: bool = False,
allow_patterns: Optional[List[str]] = None,
) -> Path:
"""
Ensures the model is available locally. If the path does not exist locally,
it is downloaded from the Hugging Face Hub.
Args:
path_or_hf_repo (str): The local path or Hugging Face repository ID of the model.
revision (str, optional): A revision id which can be a branch name, a tag, or a commit hash.
Returns:
Path: The path to the model.
"""
model_path = Path(path_or_hf_repo)
if not model_path.exists():
model_path = Path(
snapshot_download(
repo_id=path_or_hf_repo,
revision=revision,
allow_patterns=allow_patterns
or [
"*.json",
"*.jsonl",
"*.safetensors",
"*.py",
"*.model",
"*.tiktoken",
"*.txt",
"*.jinja",
],
force_download=force_download,
)
)
return model_path
def load_model(
model_path: Path,
lazy: bool = False,
**kwargs,
) -> nn.Module:
"""
Load and initialize the model from a given path.
Args:
model_path (Path): The path to load the model from.
lazy (bool): If False eval the model parameters to make sure they are
loaded in memory before returning, otherwise they will be loaded
when needed. Default: ``False``
revision (str, optional): A revision id which can be a branch name,
a tag, or a commit hash. Default: ``None``.
quantize_activations (bool, optional): If True, convert QuantizedLinear layers
to QQLinear layers for activation quantization. Only supported for models
quantized with 'nvfp4' or 'mxfp8' modes. Default: ``False``.
Returns:
nn.Module: The loaded and initialized model.
Raises:
FileNotFoundError: If the weight files (.safetensors) are not found.
ValueError: If the model class or args class are not found or cannot be instantiated.
"""
strict = kwargs.pop("strict", True)
# An expert-offload dir (mlx_vlm.moe_offload) is missing routed-expert
# keys by design; defer eval until patch_model swaps those modules, or
# their random-init resident weights get eagerly materialized -- the OOM
# this feature exists to avoid.
is_offload_dir = (model_path / "offload_index.json").exists()
if is_offload_dir:
strict = False
requested_lazy, lazy = lazy, True
config = load_config(model_path, **kwargs)
index_file = model_path / "model.safetensors.index.json"
weight_files = []
declared_keys: set = set()
if index_file.exists():
try:
with open(index_file) as f:
weight_map = json.load(f).get("weight_map", {})
declared_keys = set(weight_map)
weight_files = [
str(model_path / shard)
for shard in sorted(set(weight_map.values()))
if (model_path / shard).exists()
]
except (ValueError, OSError):
weight_files = []
declared_keys = set()
if not weight_files:
weight_files = [
wf
for wf in glob.glob(str(model_path / "*.safetensors"))
if not wf.endswith("consolidated.safetensors")
]
if not weight_files:
logging.error(f"No safetensors found in {model_path}")
message = f"""
No safetensors found in {model_path}
Create safetensors using the following code:
```
from transformers import AutoModelForCausalLM, AutoProcessor
model_id = " header_len:
raise RuntimeError(
f"Cannot reinterpret unsupported safetensors dtype in {path}: "
"patched header is larger than the original header."
)
try:
f.seek(8)
f.write(patched_header)
f.write(b" " * (header_len - len(patched_header)))
f.flush()
return mx.load(path)
finally:
f.seek(8)
f.write(original_header)
f.flush()
def sanitize_weights(model_obj, weights, config=None):
"""Helper function to sanitize weights if the model has a sanitize method"""
if hasattr(model_obj, "sanitize"):
if config is not None:
model_obj = model_obj(config)
weights = model_obj.sanitize(weights)
return weights
def update_module_configs(model_config, model_class, config, modules):
"""Updates configuration for model modules like text and vision modules.
Args:
model_config: The model configuration object that will be updated
model_class: The model class containing component config classes
config: Dictionary containing configuration parameters
modules: List of module names to update configs for (e.g. ["text", "vision"])
Returns:
The updated model_config object
"""
for config_name in modules:
config_attr = f"{config_name}_config"
if hasattr(model_config, config_attr) and config.get(config_attr) is not None:
config_class = getattr(model_class, f"{config_name.title()}Config")
setattr(
model_config, config_attr, config_class.from_dict(config[config_attr])
)
return model_config
def load(
path_or_hf_repo: str,
adapter_path: Optional[str] = None,
lazy: bool = False,
revision: Optional[str] = None,
strict: bool = True,
**kwargs,
) -> Tuple[nn.Module, ProcessorMixin]:
"""
Load the model and tokenizer from a given path or a huggingface repository.
Args:
path_or_hf_repo (str): The path or the huggingface repository to load the model from.
tokenizer_config (dict, optional): Configuration parameters specifically for the tokenizer.
Defaults to an empty dictionary.
adapter_path (str, optional): Path to the LoRA adapters. If provided, applies LoRA layers
to the model. Default: ``None``.
lazy (bool): If False eval the model parameters to make sure they are
loaded in memory before returning, otherwise they will be loaded
when needed. Default: ``False``
revision (str, optional): A revision id which can be a branch name,
a tag, or a commit hash. Default: ``None``.
strict (bool): Whether or not to raise an exception if weights don't
match. Default: ``True``.
quantize_activations (bool, optional): If True, convert QuantizedLinear layers
to QQLinear layers for activation quantization. Only supported for models
quantized with 'nvfp4' or 'mxfp8' modes. Default: ``False``.
Returns:
Tuple[nn.Module, TokenizerWrapper]: A tuple containing the loaded model and tokenizer.
Raises:
FileNotFoundError: If config file or safetensors are not found.
ValueError: If model class or args class are not found.
"""
force_download = kwargs.get("force_download", False)
model_path = get_model_path(
path_or_hf_repo, force_download=force_download, revision=revision
)
model = load_model(model_path, lazy, strict=strict, **kwargs)
if adapter_path is not None:
model = apply_lora_layers(model, adapter_path)
model.eval()
image_processor = load_image_processor(model_path, **kwargs)
# Get the eos_token_id from the model config
eos_token_id = getattr(model.config, "eos_token_id", None)
processor = load_processor(model_path, True, eos_token_ids=eos_token_id, **kwargs)
if image_processor is not None:
processor.image_processor = image_processor
return model, processor
def sharded_load(
repo,
tensor_group: Optional[mx.distributed.Group] = None,
pipeline_group: Optional[mx.distributed.Group] = None,
):
# Get model path with everything but weight safetensors
model_path = get_model_path(repo)
# Lazy load model to figure out what type of sharding we can do and which
# weights we need to download.
model = load_model(model_path, lazy=True, strict=False)
config = model.config.to_dict()
has_tensor_parallel = hasattr(model, "shard")
if tensor_group is not None and not has_tensor_parallel:
raise ValueError(
"The model does not support tensor parallelism but a tensor_group was provided"
)
if tensor_group is None and pipeline_group is None:
if has_tensor_parallel:
tensor_group = mx.distributed.init()
processor = load_processor(
model_path, True, eos_token_ids=config.get("eos_token_id", None)
)
image_processor = load_image_processor(model_path)
if image_processor is not None:
processor.image_processor = image_processor
if tensor_group is not None:
model.shard(tensor_group)
if pipeline_group is not None:
lm = model.language_model
# The underlying model (e.g. DeepseekV3Model) has PipelineMixin
inner = lm.model if hasattr(lm, "model") else lm
if not hasattr(inner, "pipeline"):
raise ValueError("The model does not support pipeline parallelism")
inner.pipeline(pipeline_group)
print("Materializing")
mx.eval(model.language_model.parameters())
model.eval()
# Synchronize processes to avoid timeout
mx.eval(mx.distributed.all_sum(mx.array(1.0), stream=mx.cpu))
return model, processor
def load_config(model_path: Union[str, Path], **kwargs) -> dict:
"""Load model configuration from a path or Hugging Face repo.
Args:
model_path: Local path or Hugging Face repo ID to load config from
**kwargs: Additional keyword arguments to pass to the config loader
Returns:
dict: Model configuration
Raises:
FileNotFoundError: If config.json is not found at the path
"""
if isinstance(model_path, str):
model_path = get_model_path(
model_path,
revision=kwargs.get("revision"),
force_download=kwargs.get("force_download", False),
)
try:
with open(model_path / "config.json", encoding="utf-8") as f:
config = json.load(f)
generation_config_file = model_path / "generation_config.json"
if generation_config_file.exists():
try:
with open(generation_config_file, encoding="utf-8") as f:
_merge_generation_config(config, json.load(f))
except json.JSONDecodeError:
pass
except FileNotFoundError as exc:
raise FileNotFoundError(f"Config not found at {model_path}") from exc
# GLiNER2.5 ships its encoder config in a sidecar directory instead of
# inline, so fold it in alongside the other config files. Raised outside the
# block above so the missing file is not reported as a missing config.json.
if "BoundaryExtractor" in (config.get("architectures") or ()):
if "encoder_config" not in config:
encoder_config_path = model_path / "encoder_config" / "config.json"
if not encoder_config_path.is_file():
raise FileNotFoundError(
f"GLiNER2.5 encoder config not found: {encoder_config_path}"
)
with open(encoder_config_path, encoding="utf-8") as f:
config["encoder_config"] = json.load(f)
return config
def load_image_processor(model_path: Union[str, Path], **kwargs) -> BaseImageProcessor:
if isinstance(model_path, str):
model_path = get_model_path(
model_path,
revision=kwargs.get("revision"),
force_download=kwargs.get("force_download", False),
)
if not kwargs:
config = load_config(model_path, trust_remote_code=True)
else:
config = load_config(model_path, **kwargs)
try:
model_class, _ = get_model_and_args(config)
except ValueError:
return None
image_processor = None
if hasattr(model_class, "ImageProcessor"):
init_signature = inspect.signature(model_class.ImageProcessor.__init__)
if "config" in init_signature.parameters:
image_processor = model_class.ImageProcessor(config=config)
else:
image_processor = model_class.ImageProcessor()
return image_processor
def load_processor(
model_path, add_detokenizer=True, eos_token_ids=None, **kwargs
) -> ProcessorMixin:
processor = AutoProcessor.from_pretrained(model_path, **kwargs)
if add_detokenizer:
detokenizer_class = load_tokenizer(model_path, return_tokenizer=False)
# Get the tokenizer object
tokenizer_obj = (
processor.tokenizer if hasattr(processor, "tokenizer") else processor
)
# Non-text models (depth, detection) have no decode(); skip detokenizer
try:
processor.detokenizer = detokenizer_class(tokenizer_obj)
except AttributeError:
return processor
# Create and assign the StoppingCriteria
criteria = StoppingCriteria(
eos_token_ids,
tokenizer_obj,
additional_eos_token_ids=getattr(processor, "additional_eos_token_ids", ()),
)
if hasattr(processor, "tokenizer"):
processor.tokenizer.stopping_criteria = criteria
else:
processor.stopping_criteria = criteria
return processor
def fetch_from_hub(
model_path: Path, lazy: bool = False, **kwargs
) -> Tuple[nn.Module, dict, ProcessorMixin]:
model = load_model(model_path, lazy, **kwargs)
config = load_config(model_path, **kwargs)
processor = load_processor(
model_path,
add_detokenizer=False,
eos_token_ids=config.get("eos_token_id", None),
**kwargs,
)
return model, config, processor
def make_shards(weights: dict, max_file_size_gb: int = MAX_FILE_SIZE_GB) -> list:
"""
Splits the weights into smaller shards.
Args:
weights (dict): Model weights.
max_file_size_gb (int): Maximum size of each shard in gigabytes.
Returns:
list: List of weight shards.
"""
max_file_size_bytes = max_file_size_gb << 30
shards = []
shard, shard_size = {}, 0
for k, v in weights.items():
if shard_size + v.nbytes > max_file_size_bytes:
shards.append(shard)
shard, shard_size = {}, 0
shard[k] = v
shard_size += v.nbytes
shards.append(shard)
return shards
def create_model_card(
path: Union[str, Path], hf_path: Optional[Union[str, Path]] = None
):
"""
Create model card for a converted MLX model.
Args:
path (Union[str, Path]): Local path to the converted model.
hf_path (Optional[Union[str, Path]]): Original Hugging Face repo id or local path used for conversion.
"""
from huggingface_hub import ModelCard, ModelCardData
if hf_path is None:
card = ModelCard.from_template(ModelCardData(language="en"))
else:
card = ModelCard.load(hf_path)
card.data.library_name = "mlx"
if card.data.pipeline_tag is None:
card.data.pipeline_tag = "image-text-to-text"
if card.data.tags is None:
card.data.tags = ["mlx"]
elif "mlx" not in card.data.tags:
card.data.tags += ["mlx"]
if hf_path is not None:
card.data.base_model = str(hf_path)
card.text = ""
card.save(Path(path) / "README.md")
def upload_to_hub(path: str, upload_repo: str):
"""
Uploads the model to Hugging Face hub.
Args:
path (str): Local path to the model.
upload_repo (str): Name of the HF repo to upload to.
"""
from huggingface_hub import HfApi, ModelCard, logging
from . import __version__
logging.set_verbosity_info()
card_path = Path(path) / "README.md"
card = ModelCard.load(card_path)
hf_path = card.data.base_model
if hf_path is not None:
provenance = f"""
This model was converted to MLX format from [`{hf_path}`](https://huggingface.co/{hf_path})
using mlx-vlm version **{__version__}**.
Refer to the [original model card](https://huggingface.co/{hf_path}) for more details on the model.
"""
else:
provenance = ""
card.text = dedent(f"""
# {upload_repo}
{provenance}
## Use with mlx
```bash
pip install -U mlx-vlm
```
```bash
python -m mlx_vlm.generate --model {upload_repo} --max-tokens 100 --temperature 0.0 --prompt "Describe this image." --image