ydy9038074's picture
Publish Modilify Mk2 Preview MLX
e4f7326 verified
Raw History Blame Contribute Delete
7.96 kB
"""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',
]))