"""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