"""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, )