"""Strict loading of the self-contained schema25 native MLX preview export.""" from __future__ import annotations import hashlib import json from dataclasses import dataclass from pathlib import Path from typing import Any import mlx.core as mx from mlx.utils import tree_flatten from transformers import AutoTokenizer from modilify_mk2.configuration_modilify_mk2 import ( ModilifyMk2Config, require_current_checkpoint_protocol, configure_generation_config, load_generation_config, ) from modilify_mk2.mlx_model import ( MLXModilifyMk2, LoRAConfig, create_mlx_text_backbone, inject_mlx_lora, ) @dataclass class MLXRuntime: model: MLXModilifyMk2 tokenizer: Any config: Any generation: Any checkpoint: Path step: int tensor_count: int PRECISION_POLICY = "gdn2_small_fp32_v1" def small_parameter_fp32(name: str) -> bool: return name.startswith("latent_deliberation.") and ( name.endswith((".dt_bias", ".a_log")) or (name.endswith(".weight") and "norm" in name.rsplit(".", 2)[-2]) ) def model_dtype(name: str, policy: str) -> str: if policy != PRECISION_POLICY: raise RuntimeError(f"Unsupported trainable precision policy: {policy}") return "float32" if small_parameter_fp32(name) else "bfloat16" def _install_sliding_encoder_masks(encoder: Any) -> None: """Align mlx-vlm's full-length masks with trimmed sliding KV chunks.""" if getattr(encoder, "_mlx_training_sliding_masks", False): return original = encoder._make_encoder_masks window = int(encoder.text_config.sliding_window) def make_masks(h, cache, attention_mask=None, mm_token_type_ids=None): masks = original(h, cache, attention_mask, mm_token_type_ids) if isinstance(masks, dict): return masks query_length = int(h.shape[1]) maximum = window - 1 + query_length return [ mask[..., -maximum:] if ( layer.layer_type == "sliding_attention" and isinstance(mask, mx.array) and mask.shape[-1] > maximum ) else mask for layer, mask in zip(encoder.decoder.layers, masks) ] encoder._make_encoder_masks = make_masks encoder._mlx_training_sliding_masks = True def load_runtime( model_path: str | Path, *, canvas_length: int | None, max_new_tokens: int, max_denoising_steps: int | None, repetition_penalty: float, commit_failure_budget: float | None = None, commit_top_k: int | None = None, commit_min_p: float | None = None, commit_target_confidence: float | None = None, ) -> MLXRuntime: root = resolve_model_path(model_path) manifest = json.loads((root / 'export_manifest.json').read_text(encoding='utf-8')) require_current_checkpoint_protocol(manifest) if manifest.get('format') != 'modilify_mk2_native_mlx_inference_v1': raise RuntimeError('Unsupported native MLX inference export format.') config = ModilifyMk2Config.from_pretrained(root, local_files_only=True) for name, value in ( ('commit_failure_budget', commit_failure_budget), ('commit_top_k', commit_top_k), ('commit_min_p', commit_min_p), ('commit_target_confidence', commit_target_confidence), ): if value is not None: setattr(config, name, value) if canvas_length is not None and not 1 <= canvas_length <= int(config.canvas_length): raise ValueError('--canvas-length must be within the exported canvas.') backbone = create_mlx_text_backbone(config) lora = dict(manifest['lora_config']) lora['target_modules'] = tuple(lora['target_modules']) inject_mlx_lora(backbone, LoRAConfig(**lora)) model = MLXModilifyMk2(backbone, config) model.latent_deliberation.set_dtype(mx.bfloat16) expected = dict(tree_flatten(model.parameters())) trainables = dict(tree_flatten(model.trainable_parameters())) if set(trainables) != set(manifest['trainable_names']): raise RuntimeError('Export trainable topology does not match this model.') specs = manifest['tensors'] index = json.loads((root / 'model.safetensors.index.json').read_text(encoding='utf-8')) weight_map = index['weight_map'] if set(expected) != set(specs) or set(expected) != set(weight_map): raise RuntimeError('Export must cover every model tensor exactly once.') shard_names = [item['file'] for item in manifest['shards']] if len(shard_names) != len(set(shard_names)) or set(shard_names) != set(weight_map.values()): raise RuntimeError('Export shard list and weight index disagree.') for name in trainables: if specs[name]['dtype'] != model_dtype(name, manifest['precision_policy']): raise RuntimeError(f'Export trainable precision mismatch: {name}') for shard in manifest['shards']: name = shard['file'] if Path(name).name != name: raise RuntimeError('Export shards must be local filenames.') path = root / name if not path.is_file() or path.stat().st_size != shard['bytes']: raise RuntimeError(f'Export shard is missing or truncated: {name}') digest = hashlib.sha256() with path.open('rb') as handle: for block in iter(lambda: handle.read(8 * 1024 * 1024), b''): digest.update(block) if digest.hexdigest() != shard['sha256']: raise RuntimeError(f'Export shard checksum mismatch: {name}') raw = mx.load(str(path)) declared = {key for key, value in weight_map.items() if value == name} if set(raw) != declared or set(shard['tensors']) != declared: raise RuntimeError(f'Export shard tensor index mismatch: {name}') for key, value in raw.items(): dtype = str(value.dtype).removeprefix('mlx.core.') if (list(value.shape) != specs[key]['shape'] or value.shape != expected[key].shape or dtype != specs[key]['dtype']): raise RuntimeError(f'Export tensor shape/dtype mismatch: {key}') if '.lora_' in key and '.experts.' not in key: # Schema25 restores dense adapters by transposing canonical # masters. Recreate that column-major layout: BF16 matmul # reductions can differ when Safetensors makes it row-major. raw[key] = mx.contiguous(value.T).T model.load_weights(list(raw.items()), strict=False) mx.eval(*raw.values()) del raw if canvas_length is not None: config.canvas_length = canvas_length model.eval() _install_sliding_encoder_masks(model.model.encoder) tokenizer = AutoTokenizer.from_pretrained(root, trust_remote_code=True, local_files_only=True) generation = configure_generation_config( load_generation_config(root), tokenizer, max_new_tokens=max_new_tokens, max_denoising_steps=max_denoising_steps, repetition_penalty=repetition_penalty, ) return MLXRuntime(model, tokenizer, config, generation, root, int(manifest['global_step']), len(trainables)) def resolve_model_path(model: str | Path) -> Path: """Resolve a local release or download its snapshot from Hugging Face.""" local = Path(model).expanduser() if local.is_dir(): if not (local / 'export_manifest.json').is_file(): raise ValueError(f'Missing export_manifest.json in model directory: {local}') return local.resolve() if local.is_absolute() or str(model).startswith(('.', '~')): raise FileNotFoundError(f'Model directory does not exist: {local}') from huggingface_hub import snapshot_download return Path(snapshot_download(repo_id=str(model), allow_patterns=[ 'export_manifest.json', 'config.json', 'generation_config.json', 'model*.safetensors', 'model.safetensors.index.json', 'tokenizer.json', 'tokenizer_config.json', 'chat_template.jinja', ]))