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 = "" model = AutoModelForCausalLM.from_pretrained(model_id) processor = AutoProcessor.from_pretrained(model_id) model.save_pretrained("") processor.save_pretrained("") ``` Then use the as the --hf-path in the convert script. ``` python -m mlx_vlm.convert --hf-path --mlx-path ``` """ raise FileNotFoundError(message) weights = {} for wf in weight_files: weights.update(_load_safetensors(wf)) model_class, _ = get_model_and_args(config=config, model_path=model_path) # Initialize text and vision configs if not present config.setdefault("text_config", config.pop("llm_config", {})) config.setdefault("vision_config", {}) config.setdefault("audio_config", {}) ple_storage = config["text_config"].get("ple_storage") if ple_storage and (manifest := ple_storage.get("manifest")): manifest_path = Path(manifest) if not manifest_path.is_absolute(): ple_storage["manifest"] = str(model_path / manifest_path) has_quantization = "quantization" in config # Initialize model config and update it with module configs model_config = model_class.ModelConfig.from_dict(config) modules = ["text", "vision", "perceiver", "projector", "audio"] model_config = update_module_configs(model_config, model_class, config, modules) model_config = apply_generation_config_defaults(model_config, config) if hasattr(model_config, "model_path"): model_config.model_path = str(model_path) model = model_class.Model(model_config) quantization_config = config.get("quantization_config", None) if quantization_config is None: quantization_config = config.get("text_config", {}).get( "quantization_config", None ) if quantization_config is not None: config["quantization_config"] = quantization_config weights, transformed_quantization = _transform_modelopt_nvfp4_weights( weights, quantization_config ) if transformed_quantization is None: weights, transformed_quantization = _transform_compressed_tensors_weights( weights, quantization_config ) if transformed_quantization is not None: config["quantization"] = transformed_quantization config["quantization_config"] = transformed_quantization if not has_quantization: quantization_config = config.get("quantization_config", None) if quantization_config is None: text_config = config.get("text_config", {}) quantization_config = text_config.get("quantization_config", None) if quantization_config is not None: config["quantization_config"] = quantization_config if quantization_config is None: hf_quant_path = model_path / "hf_quant_config.json" if hf_quant_path.exists(): with open(hf_quant_path) as f: hf_quant = json.load(f).get("quantization", {}) if hf_quant.get("quant_algo") == "NVFP4": # nvfp4 implies group_size=16 config["quantization"] = config["quantization_config"] = { "group_size": 16, "bits": 4, "mode": "nvfp4", } if quantization_config is not None: quant_method = quantization_config.get("quant_method") quantization = None if quant_method == "compressed-tensors": if quantization_config.get("format") == "mxfp4-pack-quantized": quantization = {"group_size": 32, "bits": 4, "mode": "mxfp4"} else: quantization = {"group_size": 32, "bits": 4, "mode": "affine"} elif quant_method == "mxfp4": quantization = {"group_size": 32, "bits": 4, "mode": "mxfp4"} elif quant_method == "fp8": from .fp8 import transform_fp8_weights weights, quantization = transform_fp8_weights(weights, config) # TODO: Refactor DeepSeek-V4 to use the shared FP8 transform. if quantization is None and config.get("model_type") == "deepseek_v4": from .models.deepseek_v4.language import make_quantization_config quantization = make_quantization_config(model) elif ( quant_method == "modelopt" and quantization_config.get("quant_algo") == "MXFP8" ): make_config = _language_model_quantization_config(model) if make_config is not None: quantization = make_config(model) elif quant_method in ("awq", "gptq"): logging.warning( "Quantization method %s is not supported in mlx_vlm.load_model()", quant_method, ) if quantization is not None: config["quantization"] = quantization config["quantization_config"] = quantization if has_quantization: for quantization_key in ("quantization", "quantization_config"): quantization_value = getattr(model_config, quantization_key, None) if quantization_value is not None: config[quantization_key] = quantization_value weights = sanitize_weights(model, weights) if hasattr(model_class, "VisionModel"): if hasattr(model_config, "vision_config"): weights = sanitize_weights( model_class.VisionModel, weights, model_config.vision_config ) if hasattr(model_class, "LanguageModel"): if hasattr(model_config, "text_config"): weights = sanitize_weights( model_class.LanguageModel, weights, model_config.text_config ) if hasattr(model_class, "AudioModel"): if hasattr(model_config, "audio_config"): weights = sanitize_weights( model_class.AudioModel, weights, model_config.audio_config ) if (quantization := config.get("quantization", None)) is not None: # Handle legacy models which may or may not have vision quantized. # text-only quants of unified VLM families set vision_config to null, # so coerce None to {} to avoid AttributeError on the .get below. # TODO: Re-upload the models with the new quantization config and remove this skip_vision = (config.get("vision_config") or {}).get("skip_vision", False) quantized_model = model # Stock MLX rejects bits=1; route those layers to our Metal kernel. replace_one_bit_modules(quantized_model, quantization, weights) def get_class_predicate(p, m): per_module_quantization = _quantization_for_module_path( config["quantization"], p, model ) # Skip legacy multimodal layers unless the checkpoint has quantized # tensors for this exact module. if ( skip_multimodal_module(p) and skip_vision and not _has_quantized_weights(p, weights) ): return False # Skip 1-bit layers already replaced above. module_quantization = ( per_module_quantization if per_module_quantization is not None else _quantization_for_path(config["quantization"], p) ) if module_quantization.get("bits") == 1: return False # Handle custom per-layer quantization, including aliases supplied # by the model. if per_module_quantization is not None: return per_module_quantization if not hasattr(m, "to_quantized"): return False # Skip layers not divisible by 64 if hasattr(m, "weight") and m.weight.size % 64 != 0: return False # Handle legacy models which may not have everything quantized return f"{p}.scales" in weights nn.quantize( quantized_model, group_size=quantization["group_size"], bits=quantization["bits"], mode=quantization.get("mode", "affine"), class_predicate=get_class_predicate, ) if kwargs.get("quantize_activations", False): if quantization is None: raise ValueError( "Activation quantization requires the model to be quantized first. " "Please use a quantized model with mode 'nvfp4' or 'mxfp8'." ) model = quantize_activations(model) if is_offload_dir: # strict=False above must not swallow a genuinely malformed offload # dir: verify every parameter left at random-init is actually an # expected expert-weight path, and fail loudly on anything else. from .moe_offload import PEREXPERT_RE, STACKED_FUSED_RE, STACKED_RE expected = {k for k, _ in tree_flatten(model.parameters())} missing = expected - set(weights) unexpected_missing = [ k for k in missing if not ( PEREXPERT_RE.match(k) or STACKED_RE.match(k) or STACKED_FUSED_RE.match(k) ) ] if unexpected_missing: raise ValueError( f"Offload dir {model_path} is missing {len(unexpected_missing)} " "resident parameters that aren't routed-expert weights (malformed " "repack() output?): " + ", ".join(sorted(unexpected_missing)[:10]) ) _drop_modules_without_weights(model, weights, declared_keys) model.load_weights(list(weights.items()), strict=strict) if is_offload_dir: from .moe_offload import patch_model model.moe_offload_store = patch_model( model, str(model_path), expert_cache_gb=kwargs.get("expert_cache_gb"), max_kv_size=kwargs.get("max_kv_size"), ) lazy = requested_lazy if not lazy: mx.eval(model.parameters()) model.model_path = model_path model.eval() return model def _load_safetensors(path: str) -> dict: try: return mx.load(path) except RuntimeError as e: if not any(dtype in str(e) for dtype in SAFETENSORS_DTYPE_FALLBACKS): raise load_error = e with open(path, "r+b") as f: header_len = struct.unpack(" 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 ``` """) card.save(card_path) api = HfApi() api.create_repo(repo_id=upload_repo, exist_ok=True) api.upload_folder( folder_path=path, repo_id=upload_repo, repo_type="model", ) print(f"Upload successful, go to https://huggingface.co/{upload_repo} for details.") def apply_repetition_penalty(logits: mx.array, generated_tokens: Any, penalty: float): """ Apply repetition penalty to specific logits based on the given context. Paper: https://arxiv.org/abs/1909.05858 Args: logits (mx.array): The logits produced by the language model. generated_tokens (any): A list of N previous tokens. penalty (float): The repetition penalty factor to be applied. Returns: logits (mx.array): Logits with repetition penalty applied to generated tokens. """ if len(generated_tokens) > 0: indices = mx.array([token for token in generated_tokens]) selected_logits = logits[:, indices] selected_logits = mx.where( selected_logits < 0, selected_logits * penalty, selected_logits / penalty ) logits[:, indices] = selected_logits return logits def save_weights( save_path: Union[str, Path], model: nn.Module, *, donate_weights: bool = False, ) -> None: """Save model weights into specified directory.""" if isinstance(save_path, str): save_path = Path(save_path) weights = dict(tree_flatten(model.parameters())) save_path.mkdir(parents=True, exist_ok=True) shards = make_shards(weights) shards_count = len(shards) shard_file_format = ( "model-{:05d}-of-{:05d}.safetensors" if shards_count > 1 else "model.safetensors" ) total_size = sum(v.nbytes for v in weights.values()) index_data = {"metadata": {"total_size": total_size}, "weight_map": {}} # Write the weights and make sure no references are kept other than the # necessary ones if donate_weights: model.update(tree_map(lambda _: mx.array([]), model.parameters())) weights.clear() del weights for i in range(len(shards)): shard = shards[i] shards[i] = None shard_name = shard_file_format.format(i + 1, shards_count) shard_path = save_path / shard_name mx.save_safetensors(str(shard_path), shard, metadata={"format": "mlx"}) for weight_name in shard.keys(): index_data["weight_map"][weight_name] = shard_name del shard index_data["weight_map"] = { k: index_data["weight_map"][k] for k in sorted(index_data["weight_map"]) } with open(save_path / "model.safetensors.index.json", "w") as f: json.dump( index_data, f, indent=4, ) def save_config( config: dict, config_path: Union[str, Path], ) -> None: """Save the model configuration to the ``config_path``. The final configuration will be sorted before saving for better readability. Args: config (dict): The model configuration. config_path (Union[str, Path]): Model configuration file path. """ # Clean unused keys config.pop("_name_or_path", None) config.pop("torch_dtype", None) # sort the config for better readability config = dict(sorted(config.items())) # write the updated config to the config_path (if provided) with open(config_path, "w") as fid: json.dump(config, fid, indent=4) def load_image(image_source: Union[str, Path, BytesIO, Image.Image], timeout: int = 10): """ Helper function to load an image from either a URL, file path, data URI, or BytesIO object. """ import base64 original_source = image_source try: if isinstance(image_source, Image.Image): image = image_source elif not isinstance(image_source, (str, Path, BytesIO)): raise ValueError( f"Unsupported image source type: {type(image_source).__name__}" ) else: if isinstance(image_source, str) and image_source.startswith("data:image/"): if "," not in image_source: raise ValueError( "Invalid data URI format - missing comma separator" ) _, data = image_source.split(",", 1) image_source = BytesIO(base64.b64decode(data)) if isinstance(image_source, str) and image_source.startswith( ("http://", "https://") ): with requests.get( image_source, stream=True, timeout=timeout ) as response: response.raise_for_status() image_source = BytesIO(response.content) image = Image.open(image_source) except ValueError: raise except Exception as e: raise ValueError(f"Failed to load image from {original_source}: {e}") from e image = ImageOps.exif_transpose(image) return image.convert("RGB") def resize_image(img, max_size): ratio = min(max_size[0] / img.width, max_size[1] / img.height) new_size = (int(img.width * ratio), int(img.height * ratio)) return img.resize(new_size) def process_image(img, resize_shape, image_processor): if isinstance(img, str): img = load_image(img) if hasattr(img, "mode") and img.mode != "RGB": img = img.convert("RGB") if resize_shape is not None: if isinstance(image_processor, BaseImageProcessor): # warnings (not logging) so repeated calls in a batch dedupe. warnings.warn( f"resize_shape={resize_shape} is ignored because " f"{type(image_processor).__name__} handles its own image " "sizing; use the processor's sizing options instead." ) else: img = resize_image(img, resize_shape) return img def estimate_num_image_tokens(processor, height: int, width: int, **size_overrides): """Estimate how many language-model image tokens an image will produce. Computed from the processor's own sizing math without loading or processing any pixels, so it is cheap enough to run per candidate image when sizing a prompt budget or choosing a ``max_pixels`` cap. Args: processor: A processor (or bare image processor). Wrapped processors are unwrapped via their ``image_processor`` attribute. height: Source image height in pixels. width: Source image width in pixels. **size_overrides: Optional per-call overrides forwarded to the image processor's ``num_image_tokens``, e.g. ``max_pixels=1_000_000``. Raises: NotImplementedError: If the image processor does not expose ``num_image_tokens``. Only dynamic-resolution processors support estimation; fixed-resolution models produce a constant token count regardless of image size. """ image_processor = getattr(processor, "image_processor", processor) counter = getattr(image_processor, "num_image_tokens", None) if counter is None: raise NotImplementedError( f"{type(image_processor).__name__} does not expose " "num_image_tokens; token estimation is only available for " "dynamic-resolution image processors." ) return int(counter(height, width, **size_overrides)) def read_audio(file) -> tuple: """Read an audio file using miniaudio (or ffmpeg for m4a/aac/ogg/opus). Returns (samples_float32, sample_rate) where samples is always 2D (samples, channels). """ import io as _io if isinstance(file, bytes): file = _io.BytesIO(file) # Check if ffmpeg is needed for certain formats use_ffmpeg = False if isinstance(file, (str, Path)): ext = Path(file).suffix.lstrip(".").lower() if ext in ("m4a", "aac", "ogg", "opus"): use_ffmpeg = True elif isinstance(file, _io.BytesIO): pos = file.tell() header = file.read(12) file.seek(pos) if header[4:8] == b"ftyp" or header[:4] == b"OggS": use_ffmpeg = True if use_ffmpeg: import json as _json import shutil import subprocess ffmpeg_path = shutil.which("ffmpeg") ffprobe_path = shutil.which("ffprobe") if ffmpeg_path is None: raise RuntimeError( "ffmpeg not found. Install it: brew install ffmpeg (macOS) " "or sudo apt install ffmpeg (Linux)" ) if isinstance(file, _io.BytesIO): file.seek(0) input_data = file.read() else: input_data = None # Get info via ffprobe if ffprobe_path and input_data is not None: probe = subprocess.run( [ ffprobe_path, "-v", "quiet", "-print_format", "json", "-show_streams", "-select_streams", "a:0", "-i", "pipe:0", ], input=input_data, capture_output=True, ) elif ffprobe_path: probe = subprocess.run( [ ffprobe_path, "-v", "quiet", "-print_format", "json", "-show_streams", "-select_streams", "a:0", str(file), ], capture_output=True, ) else: probe = None sample_rate, nchannels = 44100, 1 if probe and probe.returncode == 0: info = _json.loads(probe.stdout) if info.get("streams"): stream = info["streams"][0] sample_rate = int(stream.get("sample_rate", 44100)) nchannels = int(stream.get("channels", 1)) # Decode via ffmpeg to raw PCM s16le cmd = [ffmpeg_path, "-v", "quiet"] if input_data is not None: cmd += ["-i", "pipe:0"] else: cmd += ["-i", str(file)] cmd += ["-f", "s16le", "-acodec", "pcm_s16le", "-ac", str(nchannels), "pipe:1"] result = subprocess.run(cmd, input=input_data, capture_output=True) if result.returncode != 0: raise RuntimeError(f"ffmpeg decoding failed: {result.stderr.decode()}") samples = np.frombuffer(result.stdout, dtype=np.int16) else: import miniaudio if isinstance(file, (str, Path)): info = miniaudio.get_file_info(str(file)) decoded = miniaudio.decode_file( str(file), nchannels=info.nchannels, sample_rate=info.sample_rate, ) elif isinstance(file, _io.BytesIO): file.seek(0) data = file.read() # Detect format from magic bytes if data[:4] == b"RIFF": info = miniaudio.wav_get_info(data) elif data[:3] == b"ID3" or data[:2] in (b"\xff\xfb", b"\xff\xfa"): info = miniaudio.mp3_get_info(data) elif data[:4] == b"fLaC": info = miniaudio.flac_get_info(data) else: info = miniaudio.vorbis_get_info(data) decoded = miniaudio.decode( data, nchannels=info.nchannels, sample_rate=info.sample_rate, ) else: raise TypeError(f"Unsupported file type: {type(file)}") sample_rate = decoded.sample_rate nchannels = decoded.nchannels samples = np.array(decoded.samples, dtype=np.int16) # Reshape multi-channel and convert to float32 if nchannels > 1: samples = samples.reshape(-1, nchannels) audio = samples.astype(np.float32) / 32768.0 # Ensure always 2D if audio.ndim == 1: audio = audio[:, np.newaxis] return audio, sample_rate def load_audio( file, sr: int, timeout: int = 10, ): """ Helper function to load audio from either a URL, file path, or numpy array. """ if isinstance(file, np.ndarray): audio = file.astype(np.float32, copy=False) return audio.mean(axis=1) if audio.ndim > 1 else audio from mlx_audio.audio_io import read as read_audio from mlx_audio.utils import resample_audio if isinstance(file, Path): file = str(file) if isinstance(file, str) and file.startswith(("http://", "https://")): try: response = requests.get(file, stream=True, timeout=timeout) response.raise_for_status() audio, sample_rate = read_audio(BytesIO(response.content), dtype="float32") except Exception as e: raise ValueError( f"Failed to load audio from URL: {file} with error {e}" ) from e else: audio, sample_rate = read_audio(file, dtype="float32") # Downmix before resampling: read_audio returns (samples, channels), and # resample_audio works on the last axis, so resampling stereo first would # resample the channel axis and leave the samples at the source rate. audio = np.asarray(audio, dtype=np.float32) if audio.ndim > 1: audio = audio.mean(axis=1) if sample_rate != sr: audio = resample_audio(audio, sample_rate, sr) return np.asarray(audio, dtype=np.float32) @dataclass(frozen=True) class VideoSampling: """How many frames to take from a clip, and at what rate. Every field is optional so an unset one can be filled from a lower-precedence source via :meth:`merge`, with :data:`DEFAULT_VIDEO_SAMPLING` terminating the chain. """ fps: Optional[float] = None nframes: Optional[int] = None min_frames: Optional[int] = None max_frames: Optional[int] = None frame_factor: Optional[int] = None def merge(self, fallback: "VideoSampling") -> "VideoSampling": """Return a copy with every unset field taken from ``fallback``.""" return VideoSampling( **{ f.name: ( getattr(self, f.name) if getattr(self, f.name) is not None else getattr(fallback, f.name) ) for f in fields(self) } ) DEFAULT_VIDEO_SAMPLING = VideoSampling( fps=2.0, min_frames=4, max_frames=768, frame_factor=2 ) @dataclass class VideoMetadata: """What the decode step knew about a clip, carried alongside its frames. Field names mirror ``transformers.video_utils.VideoMetadata`` so a processor ported from upstream can consume this unchanged. """ total_num_frames: int fps: float frames_indices: List[int] width: Optional[int] = None height: Optional[int] = None duration: Optional[float] = None @property def timestamps(self) -> List[float]: """Seconds into the clip for each sampled frame.""" return [idx / self.fps for idx in self.frames_indices] @property def sampled_fps(self) -> float: """Frame rate actually achieved, which clamping can push well below the requested ``fps``.""" return len(self.frames_indices) / max(self.total_num_frames, 1e-6) * self.fps def load_video( video_path: str, sampling: Optional[VideoSampling] = None, frame_sampler=None, **sampling_kwargs, ) -> Tuple[np.ndarray, VideoMetadata]: """Read a video file as a (T, C, H, W) numpy array. Samples ``nframes`` frames, a count derived from ``fps``, or indices from ``frame_sampler``. Returns source-aware :class:`VideoMetadata` alongside the frames. Sampling fields may be supplied as a :class:`VideoSampling` or as loose keyword arguments; the former wins for overlapping fields. """ import cv2 if sampling_kwargs: unknown = set(sampling_kwargs) - {f.name for f in fields(VideoSampling)} if unknown: raise TypeError( f"load_video() got unexpected keyword arguments: {sorted(unknown)}" ) sampling = (sampling or VideoSampling()).merge(VideoSampling(**sampling_kwargs)) resolved = (sampling or VideoSampling()).merge(DEFAULT_VIDEO_SAMPLING) fps = resolved.fps nframes = resolved.nframes min_frames = resolved.min_frames max_frames = resolved.max_frames frame_factor = resolved.frame_factor if video_path.startswith("file://"): video_path = video_path[7:] cap = cv2.VideoCapture(video_path) if not cap.isOpened(): raise ValueError(f"Cannot open video: {video_path}") total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) video_fps = cap.get(cv2.CAP_PROP_FPS) or 1.0 width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) duration = total_frames / video_fps def _round(n): return round(n / frame_factor) * frame_factor def _floor(n): return math.floor(n / frame_factor) * frame_factor def _ceil(n): return math.ceil(n / frame_factor) * frame_factor used_frame_sampler = False if nframes is not None: n = _round(nframes) indices = np.linspace(0, total_frames - 1, n).round().astype(int) elif frame_sampler is not None: used_frame_sampler = True source_metadata = VideoMetadata( total_num_frames=total_frames, fps=video_fps, frames_indices=list(range(total_frames)), width=width, height=height, duration=duration, ) indices = np.asarray( frame_sampler(source_metadata, fps=fps, max_frames=max_frames), dtype=int ).reshape(-1) n = len(indices) else: lo = _ceil(min_frames) hi = _floor(min(max_frames, total_frames)) n = total_frames / video_fps * fps n = min(max(n, lo), hi, total_frames) n = _floor(n) indices = np.linspace(0, total_frames - 1, n).round().astype(int) if not used_frame_sampler and not (frame_factor <= n <= total_frames): cap.release() raise ValueError( f"nframes must be in [{frame_factor}, {total_frames}], got {n}." ) if n == 0 or np.any(indices < 0) or np.any(indices >= total_frames): cap.release() raise ValueError( f"Frame indices must be within a non-empty {total_frames}-frame video." ) frames = [] for idx in indices: cap.set(cv2.CAP_PROP_POS_FRAMES, idx) ret, frame = cap.read() if not ret: break frames.append(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)) cap.release() if not frames: raise ValueError("No frames read from the video.") video_np = np.transpose(np.stack(frames, axis=0), (0, 3, 1, 2)) metadata = VideoMetadata( total_num_frames=total_frames, fps=video_fps, # Truncated where a read failed part-way, so this tracks the frames # actually returned rather than the ones asked for. frames_indices=indices[: len(frames)].tolist(), height=int(video_np.shape[2]), width=int(video_np.shape[3]), duration=total_frames / video_fps, ) return video_np, metadata _VIDEO_SAMPLING_FIELDS = tuple(f.name for f in fields(VideoSampling)) def processor_video_sampling(processor) -> VideoSampling: """The frame sampling a processor asks for, if it declares any. A processor opts in either by exposing a ``video_sampling_defaults()`` hook — the way to declare a cap whose attribute is named something else, such as Gemma 4's ``num_frames`` — or by carrying matching attributes on its video processor component. """ component = getattr(processor, "video_processor", None) if component is None: return VideoSampling() for owner in (component, processor): hook = getattr(owner, "video_sampling_defaults", None) if callable(hook): declared = hook() return ( declared if isinstance(declared, VideoSampling) else VideoSampling(**declared) ) return VideoSampling( **{ name: getattr(component, name, None) for name in ("fps", "min_frames", "max_frames") } ) def resolve_video_sampling(processor, overrides: Dict[str, Any]) -> VideoSampling: """Settle how a clip gets sampled. Caller wins, then whatever the processor declares, then the library defaults. Consumes the sampling keys from ``overrides`` so they do not travel on to the processor, which only ever sees frames that were already chosen here. """ explicit = VideoSampling( **{ name: overrides.pop(name) for name in _VIDEO_SAMPLING_FIELDS if name in overrides } ) return explicit.merge(processor_video_sampling(processor)).merge( DEFAULT_VIDEO_SAMPLING ) def process_inputs( processor, prompts, images=None, audio=None, add_special_tokens=False, padding=True, padding_side="left", return_tensors="mlx", **kwargs, ): # Get the process method from the processor process_method = getattr(processor, "process", processor) parameters = inspect.signature(process_method).parameters # Prepare arguments args = { "text": prompts, "images": images, "padding": padding, "return_tensors": return_tensors, } if "padding_side" in parameters: args["padding_side"] = padding_side # Add special tokens if supported if "add_special_tokens" in parameters: args["add_special_tokens"] = add_special_tokens for param in parameters.keys(): if param in kwargs.keys(): args[param] = kwargs.get(param, None) # Add audio if provided and supported if audio is not None and len(audio) > 0: if "audio" in parameters: args["audio"] = audio elif "audios" in parameters: args["audios"] = audio else: raise ValueError( f"Processor {processor.__class__.__name__} does not support audio parameter" ) return process_method(**args) def process_inputs_with_fallback( processor, prompts, images, audio, add_special_tokens=False, return_tensors="mlx", **kwargs, ): # First attempt with specified return_tensors try: return process_inputs( processor, prompts=prompts, images=images, audio=audio, add_special_tokens=add_special_tokens, return_tensors=return_tensors, **kwargs, ) except Exception as e: raise ValueError(f"Failed to process inputs with error: {e}") def prepare_inputs( processor, images=None, audio=None, videos=None, prompts=None, image_token_index=None, resize_shape=None, add_special_tokens=False, padding=True, padding_side="left", pad_to_uniform_size=False, return_tensors="mlx", **kwargs, ): has_images = images is not None and ( not hasattr(images, "__len__") or len(images) > 0 ) has_audio = audio is not None and (not hasattr(audio, "__len__") or len(audio) > 0) has_videos = videos is not None and ( not hasattr(videos, "__len__") or len(videos) > 0 ) if not has_images and not has_audio and not has_videos: tokenizer = ( processor.tokenizer if hasattr(processor, "tokenizer") else processor ) # Ensure pad_token exists when padding text-only inputs if padding and tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token inputs = tokenizer( prompts, add_special_tokens=add_special_tokens, padding=padding, padding_side=padding_side, return_tensors=return_tensors, ) input_ids = ( inputs.input_ids if isinstance(inputs.input_ids, mx.array) else mx.array(inputs.input_ids) ) mask = ( inputs.attention_mask if isinstance(inputs.attention_mask, mx.array) else mx.array(inputs.attention_mask) ) return { "input_ids": input_ids, "attention_mask": mask, } # Process images if images is not None: if not isinstance(images, list): images = [images] image_processor = ( processor.image_processor if hasattr(processor, "image_processor") else None ) images = [process_image(img, resize_shape, image_processor) for img in images] # For batching, we need uniform image sizes. Instead of padding to the # largest image (which adds white borders that hurt accuracy), we resize # all images to the model's expected input size. if len(images) > 1 and pad_to_uniform_size: # Get target size from image processor if available target_size = None if image_processor is not None and hasattr(image_processor, "size"): size = image_processor.size if isinstance(size, tuple): target_size = size elif isinstance(size, dict): target_size = (size.get("height", 384), size.get("width", 384)) elif isinstance(size, int): target_size = (size, size) if target_size is not None: # Resize all images to the target size resized_images = [] for img in images: if img.size != ( target_size[1], target_size[0], ): # PIL uses (width, height) img = img.resize( (target_size[1], target_size[0]), Image.Resampling.BICUBIC ) resized_images.append(img) images = resized_images else: # Fallback: pad to largest size (original behavior) max_width = max(img.width for img in images) max_height = max(img.height for img in images) padded_images = [] for img in images: if img.width != max_width or img.height != max_height: padded_img = Image.new( "RGB", (max_width, max_height), (255, 255, 255) ) x_offset = (max_width - img.width) // 2 y_offset = (max_height - img.height) // 2 padded_img.paste(img, (x_offset, y_offset)) padded_images.append(padded_img) else: padded_images.append(img) images = padded_images # Process audio if audio is not None and len(audio) > 0: if not isinstance(audio, list): audio = [audio] if len(audio) > 1 and not getattr(processor, "supports_multiple_audio", False): print( "\033[33mWarning\033[0m: Single prompt with multiple audio files is not supported yet. Using the first audio file.\n" ) audio = audio[:1] feature_extractor = getattr(processor, "feature_extractor", None) sr = ( getattr(feature_extractor, "sampling_rate", 16000) if feature_extractor is not None else 16000 ) audio = [load_audio(audio_file, sr=sr) for audio_file in audio] video_fps = None supplied_video_metadata = kwargs.pop("video_metadata", None) video_metadata = None if has_videos: if not isinstance(videos, list): videos = [videos] sampling_overrides = { name: kwargs.pop(name) for name in _VIDEO_SAMPLING_FIELDS if name in kwargs } fps_values = sampling_overrides.get("fps") if isinstance(fps_values, (list, tuple)) and len(fps_values) != len(videos): raise ValueError( f"Received {len(fps_values)} fps values for {len(videos)} videos" ) if supplied_video_metadata is not None and len(supplied_video_metadata) != len( videos ): raise ValueError("Expected one video_metadata entry per video.") component = getattr(processor, "video_processor", None) frame_sampler = ( getattr(component, "sample_frames", None) if getattr(component, "sample_frames_in_loader", False) else None ) loaded, video_fps, video_metadata = [], [], [] for video_index, v in enumerate(videos): video_sampling_overrides = dict(sampling_overrides) if isinstance(fps_values, (list, tuple)): video_sampling_overrides["fps"] = fps_values[video_index] sampling = resolve_video_sampling(processor, video_sampling_overrides) if isinstance(v, (str, bytes, Path)): arr, metadata = load_video( str(v), sampling, frame_sampler=frame_sampler ) logger.info( "video %s: sampled %d of %d frames at %.2f fps " "(source %.2f fps, %.1fs)", v, len(metadata.frames_indices), metadata.total_num_frames, metadata.sampled_fps, metadata.fps, metadata.duration, ) else: # Already-decoded frames: nothing was sampled here, so report # the requested rate and describe what we were handed. arr = v metadata = VideoMetadata( total_num_frames=len(v), fps=sampling.fps, frames_indices=list(range(len(v))), ) if supplied_video_metadata is not None: metadata = supplied_video_metadata[video_index] if isinstance(metadata, dict): metadata = VideoMetadata(**metadata) loaded.append(arr) video_fps.append(metadata.sampled_fps) video_metadata.append(metadata) videos = loaded model_inputs = {} if hasattr(processor, "image_processor") and isinstance( processor.image_processor, BaseImageProcessor ): if not isinstance(prompts, list): prompts = [prompts] if processor.pad_token is None: processor.pad_token = processor.eos_token text_chunks = [ [processor(chunk).input_ids for chunk in prompt.split("")] for prompt in prompts ] # Find the maximum length for padding max_length = max( sum(len(chunk) for chunk in chunks) + 1 for chunks in text_chunks ) # Pad and create input_ids input_ids = [] for chunks in text_chunks: ids = chunks[0] + [image_token_index] + chunks[1] padding = [processor.pad_token_id] * (max_length - len(ids)) input_ids.append(mx.array(ids + padding)) model_inputs["input_ids"] = mx.array(input_ids) pixel_values = processor.image_processor.preprocess(images=images) model_inputs["pixel_values"] = mx.array(np.stack(pixel_values)) model_inputs["attention_mask"] = mx.array( [(ids != processor.pad_token_id) for ids in input_ids] ).astype(mx.int32) else: if hasattr(processor, "tokenizer") and processor.tokenizer.pad_token is None: processor.tokenizer.pad_token = processor.tokenizer.eos_token extra = {} if has_videos: extra["videos"] = videos if video_fps is not None: extra["fps"] = video_fps if video_metadata is not None: extra["video_metadata"] = video_metadata inputs = process_inputs_with_fallback( processor, images=images, audio=audio, prompts=prompts, add_special_tokens=add_special_tokens, **extra, **kwargs, ) if "images" in inputs: inputs["pixel_values"] = inputs["images"] inputs.pop("images") attention_mask = inputs.get("attention_mask") model_inputs["attention_mask"] = ( attention_mask if attention_mask is None or isinstance(attention_mask, mx.array) else mx.array(attention_mask) ) # Convert inputs to model_inputs with mx.array if present for key, value in inputs.items(): if key not in model_inputs: if value is None: model_inputs[key] = value elif isinstance(value, (str, list, mx.array)): model_inputs[key] = value else: model_inputs[key] = mx.array(value) return model_inputs def group_images_by_shape( images: List[Image.Image], disable_grouping: bool = False, ) -> Tuple[Dict[Tuple[int, int], List[Image.Image]], Dict[Tuple[int, int], List[int]]]: """ Group images by their dimensions for efficient batch processing. Images with the same dimensions can be stacked and processed together, which is much faster than processing individually (especially on GPU). Args: images: List of PIL images to group disable_grouping: If True, each image gets its own group (useful for debugging) Returns: grouped_images: Dict mapping shape -> list of images with that shape grouped_indices: Dict mapping shape -> list of original indices Example: >>> images = [img_400x300, img_800x600, img_400x300_2] >>> grouped, indices = group_images_by_shape(images) >>> grouped {(300, 400): [img_400x300, img_400x300_2], (600, 800): [img_800x600]} >>> indices {(300, 400): [0, 2], (600, 800): [1]} """ if disable_grouping: # Each image in its own group grouped_images = {} grouped_indices = {} for i, img in enumerate(images): shape = (img.height, img.width) # Make each shape unique by adding index unique_shape = (img.height, img.width, i) grouped_images[unique_shape] = [img] grouped_indices[unique_shape] = [i] return grouped_images, grouped_indices grouped_images: Dict[Tuple[int, int], List[Image.Image]] = {} grouped_indices: Dict[Tuple[int, int], List[int]] = {} for i, img in enumerate(images): shape = (img.height, img.width) if shape not in grouped_images: grouped_images[shape] = [] grouped_indices[shape] = [] grouped_images[shape].append(img) grouped_indices[shape].append(i) return grouped_images, grouped_indices def resolve_eos_token_ids(eos_token_ids, tokenizer) -> List[int]: """Union configured EOS token ids with the tokenizer's own EOS. A checkpoint's ``eos_token_id`` can disagree with the token its chat template ends turns on -- Chandra OCR 2 configures ``<|endoftext|>`` but emits ``<|im_end|>`` -- so neither source alone is enough to stop generation. """ resolved: List[int] = [] for source in ( eos_token_ids, getattr(tokenizer, "eos_token_ids", None), getattr(tokenizer, "eos_token_id", None), ): if isinstance(source, int): source = [source] if not isinstance(source, (list, tuple, set)): continue for token_id in source: if isinstance(token_id, int) and token_id not in resolved: resolved.append(token_id) return resolved class StoppingCriteria: def __init__( self, eos_token_ids: List[int], tokenizer=None, additional_eos_token_ids: Optional[List[int]] = None, ): self.tokenizer = tokenizer self.additional_eos_token_ids = list( dict.fromkeys(additional_eos_token_ids or ()) ) self.reset(eos_token_ids) def add_eos_token_ids(self, new_eos_token_ids: Union[int, List[int]] = None): """ Add new token IDs to the list of EOS token IDs. Args: new_eos_token_ids: Integer, string, or list of integers/strings representing token IDs to add. If strings are provided, they will be converted to integers if possible. """ if new_eos_token_ids is None: return if self.tokenizer is None: raise ValueError("Processor is not provided") if new_eos_token_ids is not None: if isinstance(new_eos_token_ids, (str, int)): new_eos_token_ids = [new_eos_token_ids] resolved = [] for token in new_eos_token_ids: if isinstance(token, int): resolved.append(token) elif isinstance(token, str): resolved.append( self.tokenizer.encode(" " + token, add_special_tokens=False)[-1] ) self.eos_token_ids.extend(resolved) def reset(self, eos_token_ids: List[int] = None): resolved = resolve_eos_token_ids(eos_token_ids, self.tokenizer) resolved.extend( token_id for token_id in self.additional_eos_token_ids if token_id not in resolved ) if getattr(self, "eos_token_ids", None) != resolved: self.eos_token_ids = resolved def __call__(self, input_ids: mx.array) -> bool: return input_ids in self.eos_token_ids class ThinkingBudgetCriteria: """ Enforces a budget on thinking tokens. Tracks tokens within thinking blocks (between start and end tokens) and forces a closing sequence (e.g. ``\\n``) when budget is exceeded. """ def __init__( self, tokenizer, thinking_budget: int, thinking_end_token: str = "", thinking_start_token: Optional[str] = None, enable_thinking: bool = False, prompt_preopens_thinking: bool = False, ): self.tokenizer = tokenizer self.thinking_budget = thinking_budget self.enable_thinking = enable_thinking self.prompt_preopens_thinking = prompt_preopens_thinking # Resolve token IDs from strings self.thinking_end_token_id = tokenizer.encode( thinking_end_token, add_special_tokens=False )[-1] self.thinking_start_token_id = tokenizer.encode( thinking_start_token, add_special_tokens=False )[-1] self._forced_sequence: List[int] = [] newline_ids = tokenizer.encode("\n", add_special_tokens=False) if newline_ids: self._forced_sequence.append(newline_ids[-1]) self._forced_sequence.append(self.thinking_end_token_id) self._forced_index = 0 self.in_thinking = self.enable_thinking and self.prompt_preopens_thinking self.thinking_token_count = 0 self.budget_exceeded = False self.forced_token_id = None def reset_thinking_state(self): """Reset thinking state between generations.""" self.in_thinking = self.enable_thinking and self.prompt_preopens_thinking self.thinking_token_count = 0 self.budget_exceeded = False self._forced_index = 0 def __call__(self, token_id: int) -> Optional[int]: """Process a token and return a forced token ID if budget exceeded, else None.""" if self.enable_thinking and token_id == self.thinking_start_token_id: self.in_thinking = True return None if token_id == self.thinking_end_token_id: self.in_thinking = False self.budget_exceeded = False self._forced_index = 0 return None if self.in_thinking: self.thinking_token_count += 1 if self.thinking_token_count > self.thinking_budget: self.budget_exceeded = True if self.budget_exceeded and self._forced_index < len(self._forced_sequence): forced = self._forced_sequence[self._forced_index] self._forced_index += 1 self.forced_token_id = forced return forced self.forced_token_id = None return None def pop_forced_token_id(self) -> Optional[int]: """Return and clear the pending forced token ID, if any.""" if self.forced_token_id is None or not self.enable_thinking: return None forced_token_id = self.forced_token_id self.forced_token_id = None return forced_token_id def print_array_report(t: mx.array, label: Optional[str]) -> dict: """ Return a dictionary report of an MLX array similar to PyTorch's tensor representation. Args: arr: MLX array to analyze Returns: Dictionary containing shape, dtype, value representation, and statistics """ # Get basic statistics mean_val = mx.mean(t) std_val = mx.std(t) min_val = mx.min(t) max_val = mx.max(t) report = { "shape": f"{tuple(t.shape)}", "dtype": str(t.dtype), "value": repr(t), "mean": f"array({mean_val}, dtype={t.dtype})", "std": f"array({std_val}, dtype={t.dtype})", "min": f"array({min_val}, dtype={t.dtype})", "max": f"array({max_val}, dtype={t.dtype})", "label": label if label else "array", } # Print each field, handling 'value' specially print("{") for key, value in report.items(): if key == "value": print(f" '{key}': {value},") # No quotes around value else: print(f" '{key}': {repr(value)},") print("}") return report def should_add_special_tokens(model_type: str, processor) -> bool: """Return whether tokenization should add markers outside the chat template.""" template_owns_markers = { "gemma3", "gemma3n", "gemma4", "gemma4_unified", "laguna", } if model_type not in template_owns_markers: return True return getattr(processor, "chat_template", None) is None