"""Schema25 MLX GDN2 trajectory processor, readers, and strict restoration.""" from __future__ import annotations import math from typing import Any import mlx.core as mx from mlx import nn from .mlx_state import MLXLatentState from dataclasses import replace import weakref from .mlx_gdn2_trajectory import GDN2TrajectoryMemory def fp32_attention( query: mx.array, keys: mx.array, values: mx.array, *, attn_mask: mx.array | None = None, ) -> mx.array: """Small memory-attention reductions in FP32 with model-dtype output.""" dtype = query.dtype if query.ndim == 4 and keys.ndim == 4 and keys.shape[-2] >= 64: # The working bus attends over the full 256-token canvas. Metal's # fused attention avoids materializing its FP32 score/probability grid # while retaining FP32 accumulation and gradients through rel_bias. mask = attn_mask if mask is not None: mask = (mx.where(mask, 0.0, -1.0e9) if mask.dtype == mx.bool_ else mask.astype(mx.float32)) return mx.fast.scaled_dot_product_attention( query.astype(mx.float32), keys.astype(mx.float32), values.astype(mx.float32), scale=1.0 / math.sqrt(max(query.shape[-1], 1)), mask=mask, ).astype(dtype) scores = mx.matmul(query.astype(mx.float32), mx.swapaxes(keys.astype(mx.float32), -1, -2)) scores = scores / math.sqrt(max(query.shape[-1], 1)) if attn_mask is not None: if attn_mask.dtype == mx.bool_: scores = mx.where(attn_mask, scores, -1.0e9) else: scores = scores + attn_mask.astype(mx.float32) probabilities = mx.nan_to_num(mx.softmax(scores, axis=-1), nan=0.0) return mx.matmul(probabilities, values.astype(mx.float32)).astype(dtype) class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1.0e-6) -> None: super().__init__() self.weight = mx.ones((dim,)) self.eps = eps def __call__(self, hidden: mx.array) -> mx.array: value = hidden.astype(mx.float32) scale = mx.rsqrt(mx.mean(mx.square(value), axis=-1, keepdims=True) + self.eps) return (value * scale * self.weight.astype(mx.float32)).astype(hidden.dtype) class SwiGLU(nn.Module): def __init__(self, dim: int, hidden: int) -> None: super().__init__() self.gate = nn.Linear(dim, hidden, bias=False) self.up = nn.Linear(dim, hidden, bias=False) self.down = nn.Linear(hidden, dim, bias=False) def __call__(self, hidden: mx.array) -> mx.array: return self.down(nn.silu(self.gate(hidden)) * self.up(hidden)) class RankAttention(nn.Module): def __init__(self, dim: int, num_heads: int, kv_rank: int) -> None: super().__init__() if kv_rank % num_heads: raise ValueError("Attention K/V rank must divide heads.") self.num_heads, self.kv_rank = num_heads, kv_rank self.head_dim = kv_rank // num_heads self.q_proj = nn.Linear(dim, kv_rank, bias=False) self.k_proj = nn.Linear(dim, kv_rank, bias=False) self.v_proj = nn.Linear(dim, kv_rank, bias=False) self.o_proj = nn.Linear(kv_rank, dim, bias=False) self.q_norm = RMSNorm(dim) self.k_norm = RMSNorm(dim) def __call__(self, query: mx.array, keys: mx.array, values: mx.array, attn_mask: mx.array | None = None) -> mx.array: batch, queries, _ = query.shape key_len = keys.shape[1] heads, head_dim = self.num_heads, self.head_dim reshape_q = lambda x: x.reshape(batch, queries, heads, head_dim).transpose(0, 2, 1, 3) reshape_k = lambda x: x.reshape(batch, key_len, heads, head_dim).transpose(0, 2, 1, 3) q = reshape_q(self.q_proj(self.q_norm(query))) k = reshape_k(self.k_proj(self.k_norm(keys))) v = reshape_k(self.v_proj(values)) mask = attn_mask if mask is not None and mask.ndim == 2: mask = mask[None, None, :, :] elif mask is not None and mask.ndim == 3: mask = mask[:, None] context = fp32_attention(q, k, v, attn_mask=mask) return self.o_proj(context.transpose(0, 2, 1, 3).reshape(batch, queries, self.kv_rank)) class DecoderMemoryBus(nn.Module): """Working or persistent sidecar readers on full-attention trunk layers.""" def __init__(self, hidden_size: int, num_heads: int, num_readers: int, memory_dim: int, kv_rank: int, *, relative_bias: bool = False, address_with_identity: bool = False, max_relative_span: int = 256) -> None: super().__init__() if num_readers > 0 and (kv_rank % num_heads or hidden_size % num_heads): raise ValueError("Memory bus rank and hidden size must divide heads.") self.hidden_size, self.num_heads = hidden_size, num_heads self.num_readers, self.kv_rank = num_readers, kv_rank self.head_dim = kv_rank // num_heads if num_heads else kv_rank self.relative_bias = relative_bias self.address_with_identity = address_with_identity self.max_relative_span = max_relative_span self.memory_norm = RMSNorm(memory_dim) self.address_norm = RMSNorm(memory_dim) self.memory_to_hidden = (None if memory_dim == hidden_size else nn.Linear(memory_dim, hidden_size, bias=False)) self.k_proj = nn.Linear(hidden_size, kv_rank, bias=False) self.v_proj = nn.Linear(hidden_size, kv_rank, bias=False) self.q_norm = RMSNorm(hidden_size) self.q_proj = [nn.Linear(hidden_size, kv_rank, bias=False) for _ in range(num_readers)] self.o_proj = [nn.Linear(kv_rank, hidden_size, bias=False) for _ in range(num_readers)] self.alpha = mx.zeros((max(num_readers, 1), max(num_heads, 1))) span = max(2 * max_relative_span - 1, 1) self.rel_bias = mx.zeros((max(num_heads, 1), span)) def prepare_kv(self, memory: mx.array, slot_identity: mx.array | None = None) -> tuple[mx.array, mx.array] | None: if self.num_readers <= 0: return None if self.address_with_identity: if slot_identity is None: raise ValueError("Persistent bus requires slot identity on keys.") mapped_keys = self.address_norm(memory + slot_identity) mapped_values = self.memory_norm(memory) else: mapped_keys = mapped_values = self.memory_norm(memory) if self.memory_to_hidden is not None: mapped_keys = self.memory_to_hidden(mapped_keys) mapped_values = self.memory_to_hidden(mapped_values) batch, slots, _ = mapped_keys.shape shape = (batch, slots, self.num_heads, self.head_dim) keys = self.k_proj(mapped_keys).reshape(shape).transpose(0, 2, 1, 3) values = self.v_proj(mapped_values).reshape(shape).transpose(0, 2, 1, 3) return keys, values def _relative_mask(self, queries: int, keys: int, positions: mx.array | None = None) -> mx.array | None: if not self.relative_bias: return None if positions is not None: if positions.shape[1] != queries or queries != keys: raise ValueError("Working memory positions must match both canvas axes.") relative = mx.clip(positions[:, :, None] - positions[:, None, :] + keys - 1, 0, self.rel_bias.shape[1] - 1) return mx.take(self.rel_bias, mx.stop_gradient(relative), axis=1).transpose(1, 0, 2, 3) q, k = mx.arange(queries), mx.arange(keys) relative = mx.clip(q[:, None] - k[None, :] + keys - 1, 0, self.rel_bias.shape[1] - 1) return mx.take(self.rel_bias, mx.stop_gradient(relative), axis=1) def read(self, hidden: mx.array, reader_index: int, keys: mx.array, values: mx.array, positions: mx.array | None = None, *, key_seen: mx.array | None = None) -> mx.array: batch, canvas, _ = hidden.shape heads, head_dim = self.num_heads, self.head_dim query = self.q_proj[reader_index](self.q_norm(hidden)) query = query.reshape(batch, canvas, heads, head_dim).transpose(0, 2, 1, 3) bias = self._relative_mask(canvas, keys.shape[2], positions) if bias is not None and bias.ndim == 3: bias = bias[None] if key_seen is not None: key_mask = mx.where(key_seen[:, None, None, :], 0.0, -1e9).astype(mx.float32) bias = key_mask if bias is None else bias.astype(mx.float32) + key_mask context = fp32_attention(query, keys, values, attn_mask=bias) if key_seen is not None: context = mx.where(mx.any(key_seen, axis=-1)[:, None, None, None], context, mx.zeros_like(context)) context = context * mx.tanh(self.alpha[reader_index]).astype(hidden.dtype)[None, :, None, None] context = context.transpose(0, 2, 1, 3).reshape(batch, canvas, self.kv_rank) return hidden + self.o_proj[reader_index](context) class _CanvasBlock(nn.Module): def __init__(self, width: int, heads: int, rank: int, ffn: int, window: int, global_attention: bool) -> None: super().__init__() self.norm = RMSNorm(width) self.attn = RankAttention(width, heads, rank) self.ff_norm = RMSNorm(width) self.ff = SwiGLU(width, ffn) self.window = window self.global_attention = global_attention def __call__(self, hidden: mx.array, seen: mx.array, offsets: mx.array) -> mx.array: allowed = mx.broadcast_to(seen[:, None, :], (hidden.shape[0], hidden.shape[1], hidden.shape[1])) if not self.global_attention: allowed = allowed & (mx.abs(offsets[:, :, None] - offsets[:, None, :]) < self.window) mask = mx.where(allowed, 0.0, -1e9).astype(mx.float32) normed = self.norm(hidden) hidden = hidden + self.attn(normed, normed, normed, mask) hidden = hidden + self.ff(self.ff_norm(hidden)) return mx.where(seen[..., None], hidden, mx.zeros_like(hidden)) class _WorkingBus(DecoderMemoryBus): def prepare_kv(self, memory: tuple[mx.array, mx.array], slot_identity: mx.array | None = None): del slot_identity working, seen = memory pair = super().prepare_kv(working) return None if pair is None else (*pair, seen) def read(self, hidden: mx.array, reader_index: int, keys: mx.array, values: mx.array, seen: mx.array, positions: mx.array | None = None) -> mx.array: written = super().read(hidden, reader_index, keys, values, positions, key_seen=seen) return mx.where(seen[..., None], written, hidden) class _PersistentBus(nn.Module): def __init__(self, memory: nn.Module, readers: int) -> None: super().__init__() object.__setattr__(self, "_memory_ref", weakref.ref(memory)) self.num_readers = readers self.alpha = mx.zeros((max(readers, 1), 1)) def prepare_kv(self, memory: mx.array, seen: mx.array | None = None): if self.num_readers <= 0: return None return memory, seen def read(self, hidden: mx.array, reader_index: int, memory: mx.array, seen: mx.array | None) -> mx.array: delta = self._memory_ref().read_shared(memory, hidden) if seen is not None: delta = delta * seen[..., None].astype(delta.dtype) return hidden + mx.tanh(self.alpha[reader_index]).astype(hidden.dtype) * delta class LatentDeliberationTransformer(nn.Module): def __init__(self, *, hidden_size: int, latent_dim: int, ffn_dim: int, num_layers: int, num_heads: int, local_attention_window: int, tape_probes: int, history_kv_rank: int, num_working_readers: int, num_persistent_readers: int, working_last_block_global: bool, commit_sequence_dim: int | None = None, max_canvas_length: int = 256) -> None: super().__init__() if hidden_size != latent_dim: raise ValueError("Schema25 GDN2 requires hidden_size == latent_dim.") self.hidden_size = hidden_size self.latent_dim = latent_dim self.tape_probes = tape_probes self.packet_dim = int(commit_sequence_dim or history_kv_rank) self.trajectory = GDN2TrajectoryMemory( hidden_size, probes=tape_probes, persistent_observation_dim=self.packet_dim, ) self.blocks = [ _CanvasBlock(hidden_size, num_heads, history_kv_rank, ffn_dim, local_attention_window, bool(working_last_block_global and i == num_layers - 1)) for i in range(num_layers) ] self.output_norm = RMSNorm(hidden_size) self.working_memory_bus = _WorkingBus( hidden_size, num_heads, num_working_readers, hidden_size, history_kv_rank, relative_bias=True, max_relative_span=max_canvas_length, ) self.persistent_memory_bus = _PersistentBus( self.trajectory.persistent, num_persistent_readers, ) self.experience_in = nn.Linear(hidden_size * 3, self.packet_dim, bias=False) self.reason_embed = nn.Embedding(6, self.packet_dim) def logical_seen(self, state: MLXLatentState, canvas_head: mx.array) -> mx.array: seen = state.gdn2.seen batch, canvas = seen.shape index = (canvas_head[:, None] + mx.arange(canvas)[None, :]) % canvas return mx.take_along_axis(seen, mx.stop_gradient(index), axis=1) def __call__(self, *, token_embeddings: mx.array, confidence: mx.array, entropy: mx.array, state: MLXLatentState, canvas_head: mx.array | None = None): if state.gdn2 is None: raise RuntimeError("Schema25 GDN2 state is missing.") batch, canvas, width = token_embeddings.shape if state.memory_slots.shape != state.gdn2.persistent.shape: raise ValueError("Persistent GDN2 state shape differs from memory slots.") memory = replace(state.gdn2, persistent=state.memory_slots) hidden = self.trajectory.read(memory, token_embeddings) offsets = mx.broadcast_to(mx.arange(canvas)[None, :], (batch, canvas)) if canvas_head is not None: offsets = (offsets - canvas_head[:, None]) % canvas for block in self.blocks: hidden = block(hidden, memory.seen, offsets) hidden = self.output_norm(hidden) * memory.seen[..., None].astype(hidden.dtype) next_state = replace(state, confidence=confidence.astype(mx.float32), entropy=entropy.astype(mx.float32), gdn2=memory) return hidden, next_state def observe_state(self, state: MLXLatentState, heavy: mx.array, working: mx.array, live: mx.array, head: mx.array) -> MLXLatentState: if state.gdn2 is None: raise RuntimeError("Schema25 GDN2 state is missing.") source = mx.stop_gradient(heavy) + working updated = self.trajectory.observe(state.gdn2, source, live, head) return replace(state, gdn2=updated) def commit_write(self, *, memory: mx.array, working_state: mx.array, heavy_hidden: mx.array, committed_token_embeddings: mx.array, commit_lengths: mx.array, prefix_lengths: mx.array | None = None, commit_reason: mx.array | None = None, canvas_head: mx.array | None = None, max_commit: int | None = None): del prefix_lengths batch, canvas, width = working_state.shape count = int(mx.max(commit_lengths).item()) if max_commit is None else int(max_commit) if count <= 0: zero = mx.array(0.0) return memory, {"gate_mean": zero, "gate_max": zero, "gate_gt_01": zero, "gate_gt_05": zero, "delta_norm_mean": zero} index = mx.broadcast_to(mx.arange(count)[None, :], (batch, count)) if canvas_head is not None: index = (index + canvas_head[:, None]) % canvas selected_working = mx.take_along_axis( working_state, mx.stop_gradient(mx.broadcast_to(index[..., None], (batch, count, width))), axis=1, ) selected_heavy = mx.stop_gradient(mx.take_along_axis( heavy_hidden, mx.stop_gradient(mx.broadcast_to(index[..., None], (batch, count, width))), axis=1, )) if committed_token_embeddings.shape != selected_working.shape: raise ValueError("Committed embeddings do not match the prefix.") def normalize(value): fp32 = value.astype(mx.float32) return (fp32 * mx.rsqrt(mx.mean(mx.square(fp32), axis=-1, keepdims=True) + 1.0)).astype(value.dtype) roles = tuple(normalize(value) for value in ( selected_heavy, selected_working, mx.stop_gradient(committed_token_embeddings), )) packet = self.experience_in(mx.concatenate(roles, axis=-1)) reason = mx.zeros((batch,), mx.int32) if commit_reason is None else commit_reason reason = mx.clip(mx.where((reason == 1) | (reason == 2), 5, reason), 0, 5) packet = packet + self.reason_embed(reason)[:, None] valid = mx.arange(count)[None, :] < commit_lengths[:, None] written = self.trajectory.persistent.write_sequence(memory, packet, valid) delta = written - memory zero = mx.array(0.0) return written, {"gate_mean": zero, "gate_max": zero, "gate_gt_01": zero, "gate_gt_05": zero, "delta_norm_mean": mx.mean(mx.sqrt(mx.sum(mx.square(delta), axis=(-2, -1))))} def create_mlx_latent(config: Any) -> LatentDeliberationTransformer: readers = sum(kind == "full_attention" for kind in config.text_config.layer_types) return LatentDeliberationTransformer( hidden_size=config.text_config.hidden_size, latent_dim=config.latent_dim, ffn_dim=config.latent_ffn_dim, num_layers=config.latent_num_layers, num_heads=config.latent_num_heads, local_attention_window=config.latent_local_attention_window, tape_probes=config.latent_tape_probes, history_kv_rank=config.latent_history_kv_rank, num_working_readers=readers if config.working_memory_bus else 0, num_persistent_readers=readers if config.persistent_memory_bus else 0, working_last_block_global=config.latent_working_last_block_global, commit_sequence_dim=config.commit_sequence_dim, max_canvas_length=config.canvas_length, )