# Copyright 2026 Modilify # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0 """Compile-friendly Mk1 decoder: static prefix KV + fixed-shape sliding gather.""" from __future__ import annotations from typing import Any import time import mlx.core as mx import mlx.nn as nn from mlx_vlm.models.base import scaled_dot_product_attention from mlx_vlm.models.diffusion_gemma.language import geglu def make_static_cache(encoder, max_size: int): """Fixed-capacity prefix cache so decoder_state has a static key length.""" return encoder.make_cache(max_size=max(int(max_size), 1)) def _cache_capacity(cache) -> int: first = cache[0] keys = getattr(first, "keys", None) if keys is None: return int(getattr(first, "max_size", 0) or 0) return int(keys.shape[2]) def _prefix_length(cache) -> int: offset = getattr(cache[0], "offset", 0) if isinstance(offset, mx.array): return int(mx.max(offset).item()) return int(offset) def build_decoder_masks( *, prefix_len: int, canvas_length: int, cache_capacity: int, sliding_window: int, batch_size: int = 1, ) -> tuple[mx.array, mx.array]: """Boolean SDPA masks with static shapes for full and sliding layers.""" cache_valid = mx.arange(cache_capacity) < int(prefix_len) canvas_valid = mx.ones((canvas_length,), dtype=mx.bool_) full_row = mx.concatenate([cache_valid, canvas_valid], axis=0) full = mx.broadcast_to( full_row.reshape(1, 1, 1, -1), (batch_size, 1, canvas_length, cache_capacity + canvas_length), ) window = max(int(sliding_window) - 1, 1) slide_k = min(window, cache_capacity) if slide_k < cache_capacity: raw_idx = mx.arange(slide_k) + (int(prefix_len) - slide_k) slide_valid = (raw_idx >= 0) & (raw_idx < int(prefix_len)) else: slide_valid = mx.arange(cache_capacity) < int(prefix_len) slide_row = mx.concatenate([slide_valid, canvas_valid], axis=0) slide = mx.broadcast_to( slide_row.reshape(1, 1, 1, -1), (batch_size, 1, canvas_length, slide_k + canvas_length), ) return full, slide def _gather_window(tensor: mx.array, offset: mx.array, window: int) -> mx.array: max_size = tensor.shape[2] idx = mx.clip(mx.arange(window) + offset - window, 0, max_size - 1) gather = idx.reshape(1, 1, window, 1) return mx.take_along_axis(tensor, gather, axis=2) def _decoder_attention( attn: nn.Module, x: mx.array, mask: mx.array | None, cache, offset: mx.array, ) -> mx.array: batch, length, _ = x.shape queries = attn.q_proj(x).reshape(batch, length, attn.n_heads, attn.head_dim) queries = attn.q_norm(queries).transpose(0, 2, 1, 3) queries = attn.rope(queries, offset=offset) keys = attn.k_proj(x).reshape(batch, length, attn.n_kv_heads, attn.head_dim) values = ( attn.v_proj(x).reshape(batch, length, attn.n_kv_heads, attn.head_dim) if attn.v_proj is not None else keys ) keys = attn.k_norm(keys).transpose(0, 2, 1, 3) keys = attn.rope(keys, offset=offset) values = attn.v_norm(values).transpose(0, 2, 1, 3) encoder_keys, encoder_values = cache.decoder_state if attn.is_sliding: window = max(int(attn.config.sliding_window) - 1, 1) slide_k = min(window, int(encoder_keys.shape[2])) if slide_k < int(encoder_keys.shape[2]): encoder_keys = _gather_window(encoder_keys, offset, slide_k) encoder_values = _gather_window(encoder_values, offset, slide_k) keys = mx.concatenate([encoder_keys, keys], axis=2) values = mx.concatenate([encoder_values, values], axis=2) output = scaled_dot_product_attention( queries, keys, values, cache=None, scale=attn.scale, mask=mask ) output = output.transpose(0, 2, 1, 3).reshape(batch, length, -1) return attn.o_proj(output) def _experts_unsorted(experts: nn.Module, x: mx.array, top_k_indices, top_k_weights): """Expert FFN without argsort gather, so the decoder graph can compile.""" x = mx.expand_dims(x, (-2, -3)) gate_up = experts.gate_up_proj(x, top_k_indices, sorted_indices=False) gate = gate_up[..., : experts.hidden_dims] up = gate_up[..., experts.hidden_dims :] y = experts.down_proj(geglu(gate, up), top_k_indices, sorted_indices=False) y = y.squeeze(-2) return (y * top_k_weights[..., None]).sum(axis=-2) def _decoder_layer( layer: nn.Module, x: mx.array, mask: mx.array | None, cache, offset: mx.array, ) -> mx.array: residual = x hidden = layer.input_layernorm(x) hidden = _decoder_attention(layer.self_attn, hidden, mask, cache, offset) hidden = layer.post_attention_layernorm(hidden) hidden = residual + hidden residual = hidden shared = layer.pre_feedforward_layernorm(hidden) shared = layer.mlp(shared) shared = layer.post_feedforward_layernorm_1(shared) flat = residual.reshape(-1, residual.shape[-1]) top_k_indices, top_k_weights = layer.router(flat) routed = layer.pre_feedforward_layernorm_2(flat) routed = layer.experts(routed, top_k_indices, top_k_weights) routed = routed.reshape(residual.shape) routed = layer.post_feedforward_layernorm_2(routed) hidden = layer.post_feedforward_layernorm(shared + routed) return residual + hidden def _attn_residual( layer: nn.Module, x: mx.array, mask: mx.array | None, cache, offset: mx.array, ) -> mx.array: residual = x hidden = layer.input_layernorm(x) hidden = _decoder_attention(layer.self_attn, hidden, mask, cache, offset) hidden = layer.post_attention_layernorm(hidden) return residual + hidden def _ffn_residual(layer: nn.Module, hidden: mx.array) -> mx.array: residual = hidden shared = layer.pre_feedforward_layernorm(hidden) shared = layer.mlp(shared) shared = layer.post_feedforward_layernorm_1(shared) flat = residual.reshape(-1, residual.shape[-1]) top_k_indices, top_k_weights = layer.router(flat) routed = layer.pre_feedforward_layernorm_2(flat) routed = layer.experts(routed, top_k_indices, top_k_weights) routed = routed.reshape(residual.shape) routed = layer.post_feedforward_layernorm_2(routed) hidden = layer.post_feedforward_layernorm(shared + routed) return (residual + hidden) * layer.layer_scalar def make_compiled_attn_layers(decoder: nn.Module, cache) -> list: """Compile attention residuals only. Expert FFNs stay eager.""" compiled = [] for layer, layer_cache in zip(decoder.layers, cache): def _fn(x, offset, mask, _layer=layer, _cache=layer_cache): return _attn_residual(_layer, x, mask, _cache, offset) compiled.append(mx.compile(_fn, shapeless=True)) return compiled def decoder_hidden_states( decoder: nn.Module, canvas_ids: mx.array, latent_context: mx.array, cache, offset: mx.array, full_mask: mx.array, slide_mask: mx.array, compiled_attn_layers=None, profiler=None, ) -> mx.array: hidden = decoder._embed_canvas( canvas_ids, self_conditioning_embeddings=latent_context, ) if profiler is None: for layer, layer_cache in zip(decoder.layers, cache): mask = slide_mask if layer.layer_type == "sliding_attention" else full_mask hidden = _decoder_layer(layer, hidden, mask, layer_cache, offset) hidden = hidden * layer.layer_scalar return decoder.norm(hidden) attn_acc = 0.0 moe_acc = 0.0 for layer, layer_cache in zip(decoder.layers, cache): mask = slide_mask if layer.layer_type == "sliding_attention" else full_mask mx.eval(hidden) started = time.perf_counter() hidden = _attn_residual(layer, hidden, mask, layer_cache, offset) mx.eval(hidden) attn_acc += time.perf_counter() - started started = time.perf_counter() hidden = _ffn_residual(layer, hidden) mx.eval(hidden) moe_acc += time.perf_counter() - started profiler.add("attn", attn_acc) profiler.add("moe", moe_acc) return decoder.norm(hidden) def make_compiled_decoder( decoder: nn.Module, cache, softcap: float, ): """Compile canvas + latent-context → (logits, hidden).""" def _step( canvas_ids: mx.array, latent_context: mx.array, offset: mx.array, full_mask: mx.array, slide_mask: mx.array, ): hidden = decoder_hidden_states( decoder, canvas_ids, latent_context, cache, offset, full_mask, slide_mask, ) logits = decoder.embed_tokens.as_linear(hidden) logits = mx.tanh(logits.astype(mx.float32) / softcap) * softcap return logits, hidden return mx.compile(_step, shapeless=True) def quantize_experts(model: nn.Module, *, bits: int = 8, group_size: int = 64) -> None: """Quantize MoE expert projections only. Attention / embeddings stay bf16.""" if bits >= 16: return def predicate(path: str, module: nn.Module): if "experts" not in path: return False if not hasattr(module, "to_quantized"): return False if not (path.endswith("gate_up_proj") or path.endswith("down_proj")): return False return {"group_size": group_size, "bits": int(bits), "mode": "affine"} print(f"[mk1] quantizing experts to {bits}-bit", flush=True) before = sum(arr.nbytes for _, arr in model.parameters().items()) if False else None del before quantized = {"count": 0} def counting_predicate(path: str, module: nn.Module): result = predicate(path, module) if result: quantized["count"] += 1 return result nn.quantize(model, class_predicate=counting_predicate) print(f"[mk1] quantized {quantized['count']} expert projections", flush=True) mx.eval(model.parameters())