File size: 5,534 Bytes
e4f7326
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
"""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