ydy9038074's picture
Publish Modilify Mk2 Preview MLX
e4f7326 verified
Raw History Blame Contribute Delete
5.53 kB
"""Schema25 native MLX text trunk with GDN2 trajectory memory and dual readers."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import mlx.core as mx
from mlx import nn
from .mlx_latent import create_mlx_latent
from .mlx_state import MLXLatentState
@dataclass
class MLXCanvasOutput:
heavy_hidden: mx.array
working_state: mx.array
next_latent_state: MLXLatentState
token_embeddings: mx.array
class MLXModilifyMk2(nn.Module):
"""Shared DiffusionGemma trunk plus native commit-only trajectory module."""
def __init__(self, backbone: Any, config: Any) -> None:
super().__init__()
self.model = backbone.model
self.latent_deliberation = create_mlx_latent(config)
self.config = config
def _merge_context(self, token_embeddings: mx.array,
context: mx.array) -> mx.array:
"""Frozen self-conditioning bridge with the schema25 residual cap."""
mapper = self.model.decoder.self_conditioning
normed = mapper.pre_norm(context.astype(token_embeddings.dtype))
mapped = mapper.down_proj(
nn.gelu_approx(mapper.gate_proj(normed)) * mapper.up_proj(normed)
)
mapped_fp32 = mapped.astype(mx.float32)
energy = mx.mean(mx.square(mapped_fp32), axis=-1, keepdims=True)
token_rms = mx.sqrt(mx.mean(mx.square(token_embeddings.astype(mx.float32)),
axis=-1, keepdims=True))
cap = 0.5 * token_rms
scale = cap / mx.sqrt(energy + mx.square(cap) + 1.0e-12)
combined = token_embeddings + (mapped_fp32 * scale).astype(mapped.dtype)
return mapper.post_norm(combined)
@dataclass
class LoRAConfig:
r: int
alpha: int
dropout: float
target_modules: tuple[str, ...]
expert_r: int = 8
expert_alpha: int = 8
def _make_inference_router(base: Any):
"""Preserve the checkpoint's exact routing and expert weighting."""
import mlx.core as mx
from mlx_vlm.models.diffusion_gemma.language import Router
class InferenceRouter(Router):
def __call__(self, x):
x = mx.fast.rms_norm(x, None, self.eps)
x = x * self.scale * self._root_size
scores = self.proj(x)
k = self.config.top_k_experts
indices = mx.stop_gradient(mx.argpartition(scores, kth=-k, axis=-1)[..., -k:])
weights = mx.take_along_axis(scores, indices, axis=-1)
weights = mx.softmax(weights, axis=-1, precise=True)
return indices, weights * self.per_expert_scale[indices]
router = InferenceRouter(base.config)
router.proj = base.proj
router.scale = base.scale
router.per_expert_scale = base.per_expert_scale
return router
def mlx_text_config(config: Any):
"""Translate the validated schema25 text config to mlx-vlm's MLX model."""
from mlx_vlm.models.diffusion_gemma.config import ModelConfig
payload = config.to_dict()
payload["model_type"] = "diffusion_gemma"
payload["text_config"]["model_type"] = "diffusion_gemma_text"
payload["vision_config"] = None
return ModelConfig.from_dict(payload)
def create_mlx_text_backbone(config: Any):
"""Create a text-only MLX trunk with the official DiffusionGemma topology."""
from mlx_vlm.models.diffusion_gemma.diffusion_gemma import Model
return Model(mlx_text_config(config))
def inject_mlx_lora(model: Any, config: Any) -> int:
"""Attach mlx-lm adapters to the shared text trunk and MoE experts."""
from mlx_lm.tuner.lora import LoRALinear, LoRASwitchLinear
if min(config.r, config.alpha, config.expert_r, config.expert_alpha) <= 0:
raise ValueError("MLX LoRA ranks and alphas must be positive.")
if not 0 <= config.dropout < 1:
raise ValueError("MLX LoRA dropout must be in [0, 1).")
model.freeze()
targets = set(config.target_modules)
injected = 0
for layer in model.model.decoder.layers:
layer.router = _make_inference_router(layer.router)
layer.router.freeze()
for parent in (layer.self_attn, layer.mlp):
for name, module in list(parent.named_modules()):
if "." in name or name not in targets:
continue
adapter = LoRALinear.from_base(
module,
r=config.r,
dropout=config.dropout,
scale=config.alpha / config.r,
)
adapter.lora_a = adapter.lora_a.astype(module.weight.dtype)
adapter.lora_b = adapter.lora_b.astype(module.weight.dtype)
setattr(
parent,
name,
adapter,
)
injected += 1
for name in ("gate_up_proj", "down_proj"):
module = getattr(layer.experts, name)
adapter = LoRASwitchLinear.from_base(
module,
r=config.expert_r,
dropout=config.dropout,
scale=config.expert_alpha / config.expert_r,
)
adapter.lora_a = adapter.lora_a.astype(module.weight.dtype)
adapter.lora_b = adapter.lora_b.astype(module.weight.dtype)
setattr(
layer.experts,
name,
adapter,
)
injected += 1
if not injected:
raise RuntimeError("No MLX LoRA target modules were found.")
return injected