ydy9038074's picture
Publish Modilify Mk1 MLX runtime
a066584 verified
Raw
History Blame Contribute Delete
6.39 kB
# Copyright 2026 Modilify
# SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
"""Serializable configuration for the native Modilify Mk1 MLX runtime."""
from __future__ import annotations
from dataclasses import dataclass, field
import json
from pathlib import Path
from typing import Any
from mlx_vlm.models.diffusion_gemma.config import ModelConfig, TextConfig
from mlx_vlm.models.gemma4.config import VisionConfig
from mlx_vlm.models.qwen3_vl.config import _config_kwargs, _maybe_deserialize_config
MODEL_TYPE = "modilify_mk1"
TEXT_MODEL_TYPE = "modilify_mk1_text"
def _require_model_type(payload: dict[str, Any], expected: str) -> None:
actual = payload.get("model_type")
if actual != expected:
raise ValueError(
f"Expected model_type {expected!r}, got {actual!r}. "
"This runtime only loads native Modilify Mk1 checkpoints."
)
@dataclass
class ModilifyMk1Config:
"""Mk1 inference configuration. ``model_type`` is always ``modilify_mk1``."""
text_config: TextConfig
vision_config: VisionConfig | None = None
model_type: str = MODEL_TYPE
boi_token_id: int = 255999
eoi_token_id: int = 258882
image_token_id: int = 258880
video_token_id: int | None = None
pad_token_id: int = 0
bos_token_id: int = 2
eos_token_id: list[int] | int = field(default_factory=lambda: [1, 106])
initializer_range: float = 0.02
canvas_length: int = 256
denoise_temperature: float = 0.8
commit_failure_budget: float = 0.2
jump_failure_budget: float = 2.0
jump_on_no_progress_after: int = 12
max_ponder_steps: int = 64
min_trajectory_progress: float = 0.005
latent_dim: int = 1536
latent_dropout: float = 0.0
latent_local_attention_window: int = 128
latent_memory_slots: int = 64
latent_num_heads: int = 16
latent_num_layers: int = 4
turn_end_token_id: int = 106
repetition_penalty: float = 1.0
dtype: str | None = "bfloat16"
generation_config: dict[str, Any] | None = None
def __post_init__(self) -> None:
if self.model_type != MODEL_TYPE:
raise ValueError(
f"ModilifyMk1Config.model_type must be {MODEL_TYPE!r}, "
f"got {self.model_type!r}."
)
if isinstance(self.text_config, dict):
text_payload = dict(self.text_config)
text_payload.setdefault("model_type", TEXT_MODEL_TYPE)
self.text_config = TextConfig(
**_config_kwargs(TextConfig, text_payload)
)
@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "ModilifyMk1Config":
payload = dict(payload)
_require_model_type(payload, MODEL_TYPE)
text_payload = dict(payload.get("text_config") or {})
text_payload.setdefault("model_type", TEXT_MODEL_TYPE)
payload["text_config"] = TextConfig(
**_config_kwargs(TextConfig, text_payload)
)
if payload.get("vision_config") is not None:
payload["vision_config"] = _maybe_deserialize_config(
VisionConfig, payload.get("vision_config")
)
return cls(**_config_kwargs(cls, payload))
@classmethod
def from_json(cls, path: str | Path) -> "ModilifyMk1Config":
with Path(path).open(encoding="utf-8") as handle:
return cls.from_dict(json.load(handle))
def to_dict(self) -> dict[str, Any]:
def _as_dict(value: Any) -> Any:
if hasattr(value, "to_dict"):
return value.to_dict()
if hasattr(value, "__dict__") and not isinstance(
value, (str, int, float, bool, list, dict, type(None))
):
return {
key: _as_dict(item)
for key, item in vars(value).items()
if not key.startswith("_")
}
return value
return {
"model_type": MODEL_TYPE,
"text_config": _as_dict(self.text_config),
"vision_config": None
if self.vision_config is None
else _as_dict(self.vision_config),
"boi_token_id": self.boi_token_id,
"eoi_token_id": self.eoi_token_id,
"image_token_id": self.image_token_id,
"video_token_id": self.video_token_id,
"pad_token_id": self.pad_token_id,
"bos_token_id": self.bos_token_id,
"eos_token_id": self.eos_token_id,
"initializer_range": self.initializer_range,
"canvas_length": self.canvas_length,
"denoise_temperature": self.denoise_temperature,
"commit_failure_budget": self.commit_failure_budget,
"jump_failure_budget": self.jump_failure_budget,
"jump_on_no_progress_after": self.jump_on_no_progress_after,
"max_ponder_steps": self.max_ponder_steps,
"min_trajectory_progress": self.min_trajectory_progress,
"latent_dim": self.latent_dim,
"latent_dropout": self.latent_dropout,
"latent_local_attention_window": self.latent_local_attention_window,
"latent_memory_slots": self.latent_memory_slots,
"latent_num_heads": self.latent_num_heads,
"latent_num_layers": self.latent_num_layers,
"turn_end_token_id": self.turn_end_token_id,
"repetition_penalty": self.repetition_penalty,
"dtype": self.dtype,
}
def trunk_model_config(self) -> ModelConfig:
"""Build the borrowed trunk config object used only inside this package."""
return ModelConfig(
text_config=self.text_config,
vision_config=self.vision_config,
model_type=self.model_type,
boi_token_id=self.boi_token_id,
eoi_token_id=self.eoi_token_id,
image_token_id=self.image_token_id,
video_token_id=self.video_token_id,
initializer_range=self.initializer_range,
canvas_length=self.canvas_length,
eos_token_id=self.eos_token_id,
generation_config=self.generation_config,
dtype=self.dtype,
)
@property
def vocab_size(self) -> int:
return int(self.text_config.vocab_size)
@property
def hidden_size(self) -> int:
return int(self.text_config.hidden_size)