ydy9038074's picture
Publish Modilify Mk2 Preview MLX
e4f7326 verified
Raw History Blame Contribute Delete
3.85 kB
"""MLX rolling canvas, detached history, row tape, and commit-only slot state."""
from __future__ import annotations
from dataclasses import dataclass
import mlx.core as mx
from .mlx_gdn2_trajectory import GDN2TrajectoryState
def _overwritten(lengths: mx.array, head: mx.array, canvas: int) -> mx.array:
return ((mx.arange(canvas)[None, :] - head[:, None]) % canvas) < lengths[:, None]
@dataclass(frozen=True)
class MLXLatentState:
memory_slots: mx.array
confidence: mx.array
entropy: mx.array
age: mx.array
token_changed: mx.array
confidence_delta: mx.array
entropy_delta: mx.array
ponder_steps: mx.array
stagnation_steps: mx.array
gdn2: GDN2TrajectoryState | None = None
@classmethod
def empty(cls, batch: int, canvas: int, slots: int, latent: int,
*, dtype: mx.Dtype = mx.bfloat16,
enable_gdn2: bool = False) -> "MLXLatentState":
zeros = lambda: mx.zeros((batch, canvas), mx.float32)
persistent = (mx.zeros((batch, 16, 128, 128), mx.float32)
if enable_gdn2 else None)
return cls(persistent if persistent is not None
else mx.zeros((batch, slots, latent), dtype),
zeros(), zeros(), mx.zeros((batch, canvas), mx.int32),
zeros(), zeros(), zeros(),
mx.zeros((batch,), mx.int32), mx.zeros((batch,), mx.int32),
GDN2TrajectoryState(
mx.zeros((batch, canvas, 16, 64, 64), mx.float32),
mx.zeros((batch, 16, 64, 64), mx.float32),
persistent,
mx.zeros((batch, canvas), mx.bool_),
) if enable_gdn2 else None)
def advance_ring(self, lengths: mx.array, head: mx.array,
*, entropy_fill_value: float,
reset_clocks_mask: mx.array | None = None) -> "MLXLatentState":
mask = _overwritten(lengths, head, self.confidence.shape[-1])
committed = lengths > 0 if reset_clocks_mask is None else reset_clocks_mask
return MLXLatentState(
self.memory_slots,
mx.where(mask, 0.0, self.confidence),
mx.where(mask, entropy_fill_value, self.entropy),
mx.where(mask, 0, self.age),
mx.where(mask, 0.0, self.token_changed),
mx.where(mask, 0.0, self.confidence_delta),
mx.where(mask, 0.0, self.entropy_delta),
mx.where(committed, 0, self.ponder_steps).astype(mx.int32),
mx.where(committed, 0, self.stagnation_steps).astype(mx.int32),
None if self.gdn2 is None else self.gdn2.clear_refill(lengths, head),
)
@dataclass(frozen=True)
class MLXRollingState:
canvas: mx.array
latent: MLXLatentState
head: mx.array
def advance_ring(self, commit_lengths: mx.array, tail_tokens: mx.array,
*, entropy_fill_value: float,
reset_clocks_mask: mx.array | None = None) -> "MLXRollingState":
batch, canvas = self.canvas.shape
if commit_lengths.shape != (batch,) or tail_tokens.shape != self.canvas.shape:
raise ValueError("Commit lengths or tail tokens do not match the canvas.")
overwritten = _overwritten(commit_lengths, self.head, canvas)
old_offsets = (mx.arange(canvas)[None, :] - self.head[:, None]) % canvas
tail_source = mx.take_along_axis(tail_tokens, old_offsets, axis=1)
return MLXRollingState(
canvas=mx.where(overwritten, tail_source, self.canvas),
latent=self.latent.advance_ring(
commit_lengths, self.head, entropy_fill_value=entropy_fill_value,
reset_clocks_mask=reset_clocks_mask,
),
head=(self.head + commit_lengths) % canvas,
)