ydy9038074's picture
Publish Modilify Mk1 MLX runtime
a066584 verified
Raw
History Blame Contribute Delete
15.8 kB
# Copyright 2026 Modilify
# SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
"""Native MLX Modilify Mk1 model."""
from __future__ import annotations
from dataclasses import dataclass, replace
from pathlib import Path
from typing import Any
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_flatten
from .config import MODEL_TYPE, ModilifyMk1Config
from .convert_utils import remap_state_dict
from .fast_decode import (
decoder_hidden_states,
make_compiled_attn_layers,
quantize_experts,
)
from .language import build_mk1_backbone
from .latent_deliberation import LatentDeliberationState, LatentDeliberationTransformer
@dataclass
class ModilifyMk1StepOutput:
"""One heavy-denoise step over the rolling canvas."""
logits: mx.array | None
heavy_hidden_state: mx.array
next_latent_state: LatentDeliberationState
cache: Any
latent_context: mx.array
proposal: mx.array | None = None
proposal_confidence: mx.array | None = None
token_entropy: mx.array | None = None
greedy_proposal: mx.array | None = None
greedy_confidence: mx.array | None = None
def _softcap(logits: mx.array, cap: float) -> mx.array:
return mx.tanh(logits.astype(mx.float32) / cap) * cap
def _softmax_statistics(
logits: mx.array,
temperature: float,
) -> tuple[mx.array, mx.array, mx.array, mx.array, mx.array]:
scores = logits.astype(mx.float32) / temperature
probabilities = mx.softmax(scores, axis=-1, precise=True)
greedy_proposal = mx.argmax(probabilities, axis=-1)
greedy_confidence = mx.squeeze(
mx.take_along_axis(probabilities, greedy_proposal[..., None], axis=-1),
axis=-1,
)
token_entropy = -mx.sum(
probabilities * mx.log(mx.maximum(probabilities, 1.0e-30)),
axis=-1,
)
return scores, probabilities, greedy_proposal, greedy_confidence, token_entropy
_compiled_softmax_statistics = mx.compile(_softmax_statistics, shapeless=True)
class ModilifyMk1ForBlockDiffusion(nn.Module):
"""Inference-only multimodal Modilify Mk1 model."""
def __init__(self, config: ModilifyMk1Config) -> None:
super().__init__()
if config.model_type != MODEL_TYPE:
raise ValueError(
f"Refusing to construct Mk1 with model_type={config.model_type!r}."
)
self.config = config
self.model = build_mk1_backbone(config.trunk_model_config())
self.latent_deliberation = LatentDeliberationTransformer(
hidden_size=config.hidden_size,
latent_dim=config.latent_dim,
memory_slots=config.latent_memory_slots,
num_layers=config.latent_num_layers,
num_heads=config.latent_num_heads,
local_attention_window=config.latent_local_attention_window,
dropout=config.latent_dropout,
)
self.final_logit_softcapping = float(
config.text_config.final_logit_softcapping
)
self._decoder_compile_failed = False
self._compiled_attn_layers = None
def make_cache(self, max_size: int | None = None):
return self.model.encoder.make_cache(max_size=max_size)
def embed_canvas_tokens(self, decoder_input_ids: mx.array) -> mx.array:
return (
self.model.decoder.embed_tokens(decoder_input_ids)
* self.model.decoder.embed_scale
)
def prefill(
self,
input_ids: mx.array,
*,
attention_mask: mx.array | None = None,
cache=None,
pixel_values: mx.array | None = None,
mm_token_type_ids: mx.array | None = None,
):
if cache is None:
cache = self.make_cache()
_, cache = self.model.encoder(
input_ids,
attention_mask=attention_mask,
cache=cache,
pixel_values=pixel_values,
mm_token_type_ids=mm_token_type_ids,
)
return cache
def update_cache(self, input_ids: mx.array, *, cache, attention_mask=None):
_, cache = self.model.encoder(
input_ids,
attention_mask=attention_mask,
cache=cache,
)
return cache
def _prepare_latent_context(
self,
decoder_input_ids: mx.array,
*,
history_hidden_state: mx.array | None,
confidence: mx.array | None,
entropy: mx.array | None,
age: mx.array | None,
latent_state: LatentDeliberationState | None,
dtype: mx.Dtype,
) -> tuple[mx.array, LatentDeliberationState]:
batch_size, canvas_length = decoder_input_ids.shape
if latent_state is None:
latent_state = LatentDeliberationState.empty(
batch_size=batch_size,
canvas_length=canvas_length,
latent_dim=self.config.latent_dim,
memory_slots=self.config.latent_memory_slots,
dtype=dtype,
)
if confidence is None:
confidence = latent_state.confidence
else:
confidence = mx.squeeze(confidence.astype(mx.float32), axis=-1) if (
confidence.ndim == 3
) else confidence.astype(mx.float32)
if entropy is None:
entropy = latent_state.entropy
else:
entropy = mx.squeeze(entropy.astype(mx.float32), axis=-1) if (
entropy.ndim == 3
) else entropy.astype(mx.float32)
if age is not None:
latent_state = replace(latent_state, age=age.astype(mx.int32))
token_embeddings = self.embed_canvas_tokens(decoder_input_ids)
history = (
mx.zeros_like(token_embeddings)
if history_hidden_state is None
else history_hidden_state
)
return self.latent_deliberation(
heavy_hidden=history,
token_embeddings=token_embeddings,
confidence=confidence,
entropy=entropy,
state=latent_state,
)
def _proposal_statistics(
self,
logits: mx.array,
*,
denoise_temperature: float | None = None,
repetition_token_mask: mx.array | None = None,
repetition_penalty: float = 1.0,
) -> tuple[mx.array, mx.array, mx.array, mx.array, mx.array]:
temperature = (
self.config.denoise_temperature
if denoise_temperature is None
else float(denoise_temperature)
)
if temperature <= 0:
raise ValueError("`denoise_temperature` must be positive.")
if (
repetition_token_mask is not None
and repetition_penalty != 1.0
and repetition_penalty > 0
):
scores = logits.astype(mx.float32)
penalized = mx.where(
scores < 0,
scores * repetition_penalty,
scores / repetition_penalty,
)
mask = repetition_token_mask.astype(mx.bool_)[:, None, :]
logits = mx.where(mask, penalized, scores)
try:
scores, probabilities, greedy_proposal, greedy_confidence, token_entropy = (
_compiled_softmax_statistics(logits, temperature)
)
except ValueError:
scores, probabilities, greedy_proposal, greedy_confidence, token_entropy = (
_softmax_statistics(logits, temperature)
)
proposal = mx.random.categorical(scores, axis=-1)
proposal_confidence = mx.squeeze(
mx.take_along_axis(probabilities, proposal[..., None], axis=-1),
axis=-1,
)
return (
proposal,
proposal_confidence,
token_entropy,
greedy_proposal,
greedy_confidence,
)
def compile_attention(self, cache) -> None:
"""Compile per-layer attention residuals. Expert FFNs stay eager."""
from .fast_decode import _cache_capacity, build_decoder_masks
print("[mk1] compiling attention layers", flush=True)
compiled = make_compiled_attn_layers(self.model.decoder, cache)
prefix_len = int(getattr(cache[0], "offset", 0))
canvas = int(self.config.canvas_length)
hidden = int(self.config.hidden_size)
dtype = self.model.decoder.embed_tokens.weight.dtype
dummy = mx.zeros((1, canvas, hidden), dtype=dtype)
offset = mx.array(prefix_len)
full_mask, slide_mask = build_decoder_masks(
prefix_len=prefix_len,
canvas_length=canvas,
cache_capacity=max(_cache_capacity(cache), 1),
sliding_window=int(self.config.text_config.sliding_window),
)
try:
for layer, attn_fn in zip(self.model.decoder.layers, compiled):
mask = (
slide_mask
if layer.layer_type == "sliding_attention"
else full_mask
)
dummy = attn_fn(dummy, offset, mask)
mx.eval(dummy)
self._compiled_attn_layers = compiled
print("[mk1] attention compile ready", flush=True)
except ValueError as exc:
print(f"[mk1] attention compile fallback: {exc}", flush=True)
self._compiled_attn_layers = None
def decoder_logits(
self,
decoder_input_ids: mx.array,
latent_context: mx.array,
cache,
offset: mx.array,
full_mask: mx.array,
slide_mask: mx.array,
compiled_decoder_step=None,
) -> tuple[mx.array, mx.array]:
del compiled_decoder_step
hidden_states = decoder_hidden_states(
self.model.decoder,
decoder_input_ids,
latent_context,
cache,
offset,
full_mask,
slide_mask,
compiled_attn_layers=self._compiled_attn_layers,
)
logits = self.model.decoder.embed_tokens.as_linear(hidden_states)
return _softcap(logits, self.final_logit_softcapping), hidden_states
def __call__(
self,
*,
decoder_input_ids: mx.array,
cache,
previous_confidence: mx.array | None = None,
previous_entropy: mx.array | None = None,
token_age: mx.array | None = None,
latent_state: LatentDeliberationState | None = None,
history_hidden_state: mx.array | None = None,
decoder_attention_mask: mx.array | None = None,
return_proposal_statistics: bool = False,
denoise_temperature: float | None = None,
repetition_token_mask: mx.array | None = None,
repetition_penalty: float = 1.0,
compiled_decoder_step=None,
profiler=None,
) -> ModilifyMk1StepOutput:
"""Run one inference step over a noisy diffusion canvas."""
del compiled_decoder_step
dtype = self.model.decoder.embed_tokens.weight.dtype
if profiler is not None:
span = profiler.measure("latent", decoder_input_ids)
latent_context, next_state = self._prepare_latent_context(
decoder_input_ids,
history_hidden_state=history_hidden_state,
confidence=previous_confidence,
entropy=previous_entropy,
age=token_age,
latent_state=latent_state,
dtype=dtype,
)
if profiler is not None:
span.done(latent_context, next_state.token_latents, next_state.memory_slots)
del decoder_attention_mask
from .fast_decode import _cache_capacity, build_decoder_masks
prefix_len = int(getattr(cache[0], "offset", 0))
offset = mx.array(prefix_len)
canvas_length = int(decoder_input_ids.shape[1])
full_mask, slide_mask = build_decoder_masks(
prefix_len=prefix_len,
canvas_length=canvas_length,
cache_capacity=max(_cache_capacity(cache), 1),
sliding_window=int(self.config.text_config.sliding_window),
batch_size=int(decoder_input_ids.shape[0]),
)
hidden_states = decoder_hidden_states(
self.model.decoder,
decoder_input_ids,
latent_context,
cache,
offset,
full_mask,
slide_mask,
compiled_attn_layers=None,
profiler=profiler,
)
if profiler is not None:
span = profiler.measure("lm_head", hidden_states)
logits = self.model.decoder.embed_tokens.as_linear(hidden_states)
logits = _softcap(logits, self.final_logit_softcapping)
if profiler is not None:
span.done(logits)
statistics = (None, None, None, None, None)
if return_proposal_statistics:
if profiler is not None:
span = profiler.measure("softmax", logits)
statistics = self._proposal_statistics(
logits,
denoise_temperature=denoise_temperature,
repetition_token_mask=repetition_token_mask,
repetition_penalty=repetition_penalty,
)
if profiler is not None:
span.done(
statistics[0],
statistics[1],
statistics[2],
statistics[3],
statistics[4],
)
return ModilifyMk1StepOutput(
logits=None if return_proposal_statistics else logits,
heavy_hidden_state=hidden_states,
next_latent_state=next_state,
cache=cache,
latent_context=latent_context,
proposal=statistics[0],
proposal_confidence=statistics[1],
token_entropy=statistics[2],
greedy_proposal=statistics[3],
greedy_confidence=statistics[4],
)
def _load_weight_files(model_path: Path) -> dict[str, mx.array]:
weight_files = sorted(model_path.glob("*.safetensors"))
if not weight_files:
raise FileNotFoundError(f"No safetensors found in {model_path}")
weights: dict[str, mx.array] = {}
for weight_file in weight_files:
weights.update(mx.load(str(weight_file)))
return weights
def load(
model_path: str | Path,
*,
lazy: bool = False,
expert_bits: int = 16,
expert_group_size: int = 64,
):
"""Load a native ``modilify_mk1`` MLX checkpoint."""
model_path = Path(model_path)
config = ModilifyMk1Config.from_json(model_path / "config.json")
if config.model_type != MODEL_TYPE:
raise ValueError(
f"Refusing to load model_type={config.model_type!r}; "
f"expected {MODEL_TYPE!r}."
)
print("[mk1] constructing graph", flush=True)
model = ModilifyMk1ForBlockDiffusion(config)
print("[mk1] reading shards", flush=True)
weights = _load_weight_files(model_path)
remapped = remap_state_dict(weights.items())
if len(remapped) != len(weights) or any(key not in remapped for key in weights):
print(
f"[mk1] remapped {len(weights)} source tensors -> {len(remapped)} MLX tensors",
flush=True,
)
del weights
print(f"[mk1] loading {len(remapped)} tensors", flush=True)
model.load_weights(list(remapped.items()), strict=True)
del remapped
if not lazy:
print("[mk1] evaluating parameters", flush=True)
mx.eval(model.parameters())
if expert_bits and expert_bits < 16:
quantize_experts(
model, bits=int(expert_bits), group_size=int(expert_group_size)
)
return model, config
def parameter_names(model: nn.Module) -> list[str]:
return [name for name, _ in tree_flatten(model.parameters())]