ydy9038074's picture
Publish Modilify Mk2 Preview MLX
e4f7326 verified
Raw History Blame Contribute Delete
19.1 kB
"""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,
)