"""Schema25 Torch GDN2 trajectory state, processor, and memory readers.""" from __future__ import annotations from collections.abc import Sequence from dataclasses import dataclass, replace import weakref import math import torch from torch import nn from torch.nn import functional as F from .gdn2_trajectory import GDN2TrajectoryMemory, GDN2TrajectoryState COMMIT_REASON_NONE = 0 COMMIT_REASON_NORMAL = 1 COMMIT_REASON_FORCED_JUMP = 2 COMMIT_REASON_TERMINAL = 3 def _sdpa_mask_value(dtype: torch.dtype) -> float: """Additive SDPA mask that stays finite on MPS fp16/bf16.""" if dtype in (torch.float16, torch.bfloat16): return -1.0e4 return -1.0e9 def _fp32_scaled_dot_product_attention( query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, *, attn_mask: torch.Tensor | None = None, ) -> torch.Tensor: """Run latent-memory attention reductions in FP32, then restore dtype. These attention maps are small compared with the frozen decoder, while their outputs feed recurrent trajectory and persistent-memory paths. A BF16 reduction error therefore compounds across denoise/commit steps and is much more expensive than the modest FP32 workspace. """ output_dtype = query.dtype stable_mask = attn_mask if stable_mask is not None and stable_mask.is_floating_point(): stable_mask = stable_mask.float() query_fp32 = query.float() key_fp32 = key.float() use_batched_mm = query.ndim == 4 and key.ndim == 4 and value.ndim == 4 if use_batched_mm: query_length = int(query.shape[-2]) key_length = int(key.shape[-2]) scores = torch.bmm( query_fp32.reshape(-1, query_length, query.shape[-1]), key_fp32.reshape(-1, key_length, key.shape[-1]).transpose(1, 2), ).reshape(*query.shape[:-2], query_length, key_length) else: scores = torch.matmul(query_fp32, key_fp32.transpose(-2, -1)) scores = scores / math.sqrt(max(query.shape[-1], 1)) if stable_mask is not None: if stable_mask.dtype == torch.bool: scores = scores.masked_fill(~stable_mask, _sdpa_mask_value(torch.float32)) else: scores = scores + stable_mask probabilities = torch.softmax(scores, dim=-1) # All-masked rows are uniform under a finite mask, but keep a NaN # barrier for any remaining -inf path that MPS softmax cannot invert. probabilities = torch.nan_to_num(probabilities, nan=0.0) if use_batched_mm: output = torch.bmm( probabilities.reshape(-1, query.shape[-2], key.shape[-2]), value.float().reshape(-1, key.shape[-2], value.shape[-1]), ).reshape(*query.shape[:-2], query.shape[-2], value.shape[-1]) else: output = torch.matmul(probabilities, value.float()) return output.to(dtype=output_dtype) @dataclass class LatentDeliberationState: """Persistent slots plus per-canvas trajectory clocks. No token latents.""" memory_slots: torch.Tensor confidence: torch.Tensor entropy: torch.Tensor ponder_steps: torch.Tensor stagnation_steps: torch.Tensor gdn2: GDN2TrajectoryState @classmethod def empty( cls, *, batch_size: int, canvas_length: int, device: torch.device, ) -> "LatentDeliberationState": persistent = torch.zeros(batch_size, 16, 128, 128, device=device, dtype=torch.float32) return cls( memory_slots=persistent, confidence=torch.zeros( batch_size, canvas_length, device=device, dtype=torch.float32 ), entropy=torch.zeros( batch_size, canvas_length, device=device, dtype=torch.float32 ), ponder_steps=torch.zeros(batch_size, device=device, dtype=torch.int32), stagnation_steps=torch.zeros(batch_size, device=device, dtype=torch.int32), gdn2=GDN2TrajectoryState( cells=torch.zeros(batch_size, canvas_length, 16, 64, 64, device=device, dtype=torch.float32), row=torch.zeros(batch_size, 16, 64, 64, device=device, dtype=torch.float32), persistent=persistent, seen=torch.zeros(batch_size, canvas_length, device=device, dtype=torch.bool), ), ) @dataclass class LatentProcessorOutput: context: torch.Tensor state: LatentDeliberationState def advance_trajectory_clocks( ponder_steps: torch.Tensor, stagnation_steps: torch.Tensor, *, commit_lengths: torch.LongTensor, active_rows: torch.BoolTensor, ) -> tuple[torch.IntTensor, torch.IntTensor]: """Advance useful-ponder and stagnation clocks for each row.""" if not ( ponder_steps.shape == stagnation_steps.shape == commit_lengths.shape == active_rows.shape ): raise ValueError("Trajectory clock inputs must share shape [batch].") committed = commit_lengths.gt(0) waiting = active_rows & ~committed next_ponder = torch.where( committed, torch.zeros_like(ponder_steps), ponder_steps + waiting.to(torch.int32) ) next_stagnation = torch.where( committed, torch.zeros_like(stagnation_steps), stagnation_steps + waiting.to(torch.int32), ) return next_ponder.to(torch.int32), next_stagnation.to(torch.int32) def should_force_trajectory_jump( stagnation_steps: torch.Tensor, *, progress_scores: torch.Tensor | None = None, min_progress: float = 0.0, stagnation_threshold: int, ponder_steps: torch.Tensor | None = None, max_ponder_steps: int | None = None, ) -> torch.BoolTensor: jump = stagnation_steps.ge(stagnation_threshold) if progress_scores is not None: jump = jump & progress_scores.le(float(min_progress)) if ponder_steps is not None and max_ponder_steps is not None and max_ponder_steps > 0: jump = jump | ponder_steps.ge(max_ponder_steps) return jump.to(torch.bool) class _RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1.0e-6) -> None: super().__init__() self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, hidden: torch.Tensor) -> torch.Tensor: rms = hidden.float().square().mean(dim=-1, keepdim=True).add(self.eps).rsqrt() return (hidden.float() * rms * self.weight.float()).to(dtype=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 forward(self, hidden: torch.Tensor) -> torch.Tensor: return self.down(F.silu(self.gate(hidden)) * self.up(hidden)) class _RankAttention(nn.Module): """Sequence attention in a rank-``kv_rank`` subspace, then map back to ``dim``.""" def __init__(self, dim: int, num_heads: int, kv_rank: int) -> None: super().__init__() if kv_rank % num_heads: raise ValueError("`kv_rank` must be divisible by `num_heads`.") self.num_heads = num_heads self.kv_rank = 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 forward( self, query: torch.Tensor, keys: torch.Tensor, values: torch.Tensor, attn_mask: torch.Tensor | None = None, ) -> torch.Tensor: batch, queries, _dim = query.shape key_len = keys.shape[1] heads = self.num_heads head_dim = self.head_dim query = self.q_proj(self.q_norm(query)).view(batch, queries, heads, head_dim).transpose(1, 2) keys = self.k_proj(self.k_norm(keys)).view(batch, key_len, heads, head_dim).transpose(1, 2) values = self.v_proj(values).view(batch, key_len, heads, head_dim).transpose(1, 2) mask = attn_mask if mask is not None and mask.ndim == 2: mask = mask.view(1, 1, queries, key_len) elif mask is not None and mask.ndim == 3: mask = mask.unsqueeze(1) context = _fp32_scaled_dot_product_attention( query, keys, values, attn_mask=mask ) context = context.transpose(1, 2).reshape(batch, queries, self.kv_rank) return self.o_proj(context) class DecoderMemoryBus(nn.Module): """Read memory through a per-head gated residual.""" 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: if kv_rank % num_heads: raise ValueError("Memory bus rank must be divisible by heads.") if hidden_size % num_heads: raise ValueError("Memory bus hidden size must be divisible by heads.") self.hidden_size = hidden_size self.num_heads = num_heads self.num_readers = num_readers self.kv_rank = 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.memory_norm = _RMSNorm(memory_dim) self.address_norm = _RMSNorm(memory_dim) self.memory_to_hidden = ( nn.Identity() 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.ModuleList( [nn.Linear(hidden_size, kv_rank, bias=False) for _ in range(num_readers)] ) self.o_proj = nn.ModuleList( [nn.Linear(kv_rank, hidden_size, bias=False) for _ in range(num_readers)] ) self.alpha = nn.Parameter(torch.zeros(max(num_readers, 1), max(num_heads, 1))) span = max(2 * max_relative_span - 1, 1) self.rel_bias = nn.Parameter(torch.zeros(max(num_heads, 1), span)) self.max_relative_span = max_relative_span def prepare_kv( self, memory: torch.Tensor, slot_identity: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor] | 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.memory_to_hidden(self.address_norm(memory + slot_identity)) mapped_values = self.memory_to_hidden(self.memory_norm(memory)) else: mapped_keys = mapped_values = self.memory_to_hidden(self.memory_norm(memory)) batch, slots, _dim = mapped_keys.shape heads = self.num_heads head_dim = self.head_dim keys = self.k_proj(mapped_keys).view(batch, slots, heads, head_dim).transpose(1, 2) values = self.v_proj(mapped_values).view(batch, slots, heads, head_dim).transpose(1, 2) return keys, values def _relative_mask( self, queries: int, keys: int, device: torch.device, dtype: torch.dtype, positions: torch.Tensor | None = None, ) -> torch.Tensor | 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 = ( positions[:, :, None] - positions[:, None, :] + (keys - 1) ).clamp(0, self.rel_bias.shape[1] - 1) return self.rel_bias[:, relative].permute(1, 0, 2, 3).to(dtype=dtype) q = torch.arange(queries, device=device) k = torch.arange(keys, device=device) rel = (q[:, None] - k[None, :] + (keys - 1)).clamp(0, self.rel_bias.shape[1] - 1) return self.rel_bias[:, rel].to(dtype=dtype) def read( self, hidden: torch.Tensor, reader_index: int, keys: torch.Tensor, values: torch.Tensor, positions: torch.Tensor | None = None, *, key_seen: torch.Tensor | None = None, ) -> torch.Tensor: batch, canvas, _dim = hidden.shape heads = self.num_heads head_dim = self.head_dim query = self.q_proj[reader_index](self.q_norm(hidden)) query = query.view(batch, canvas, heads, head_dim).transpose(1, 2) bias = self._relative_mask( canvas, keys.shape[2], hidden.device, query.dtype, positions ) if bias is not None and bias.ndim == 3: bias = bias.unsqueeze(0) if key_seen is not None: key_mask = torch.zeros((batch, 1, 1, keys.shape[2]), device=hidden.device, dtype=torch.float32) key_mask = key_mask.masked_fill(~key_seen[:, None, None, :], -1e9) bias = key_mask if bias is None else bias.float() + key_mask context = _fp32_scaled_dot_product_attention( query, keys, values, attn_mask=bias ) if key_seen is not None: context = torch.where(key_seen.any(-1)[:, None, None, None], context, torch.zeros_like(context)) scale = torch.tanh(self.alpha[reader_index]).to(dtype=hidden.dtype).view(1, heads, 1, 1) context = context * scale context = context.transpose(1, 2).reshape(batch, canvas, self.kv_rank) return hidden + self.o_proj[reader_index](context) @torch.no_grad() def reset_identity_parameters(self) -> None: self.alpha.zero_() self.rel_bias.zero_() def slice_latent_state( state: LatentDeliberationState, rows: slice | torch.Tensor ) -> LatentDeliberationState: return LatentDeliberationState( memory_slots=state.memory_slots[rows], confidence=state.confidence[rows], entropy=state.entropy[rows], ponder_steps=state.ponder_steps[rows], stagnation_steps=state.stagnation_steps[rows], gdn2=GDN2TrajectoryState( cells=state.gdn2.cells[rows], row=state.gdn2.row[rows], persistent=state.gdn2.persistent[rows], seen=state.gdn2.seen[rows], ), ) def cat_latent_states(states: Sequence[LatentDeliberationState]) -> LatentDeliberationState: def cat(name: str) -> torch.Tensor: return torch.cat([getattr(state, name) for state in states], dim=0) return LatentDeliberationState( memory_slots=cat("memory_slots"), confidence=cat("confidence"), entropy=cat("entropy"), ponder_steps=cat("ponder_steps"), stagnation_steps=cat("stagnation_steps"), gdn2=GDN2TrajectoryState( cells=torch.cat([state.gdn2.cells for state in states], dim=0), row=torch.cat([state.gdn2.row for state in states], dim=0), persistent=torch.cat([state.gdn2.persistent for state in states], dim=0), seen=torch.cat([state.gdn2.seen for state in states], dim=0), ), ) def infer_commit_reason( commit_lengths: torch.Tensor, *, jump_rows: torch.Tensor | None = None, commit_token_ids: torch.Tensor | None = None, terminal_token_ids: Sequence[int] = (), ) -> torch.Tensor: """Return per-row commit-reason codes. No hard skip; writer sees the label.""" reasons = torch.full( commit_lengths.shape, COMMIT_REASON_NONE, device=commit_lengths.device, dtype=torch.long, ) committed = commit_lengths.gt(0) default = ( COMMIT_REASON_NORMAL ) reasons = torch.where(committed, torch.full_like(reasons, default), reasons) if jump_rows is not None: reasons = torch.where( committed & jump_rows.to(dtype=torch.bool), torch.full_like(reasons, COMMIT_REASON_FORCED_JUMP), reasons, ) if commit_token_ids is not None and terminal_token_ids: positions = torch.arange( commit_token_ids.shape[1], device=commit_token_ids.device )[None, :] selected = positions.lt(commit_lengths[:, None]) terminal = torch.zeros_like(committed) for token_id in terminal_token_ids: terminal |= (commit_token_ids.eq(int(token_id)) & selected).any(dim=-1) reasons = torch.where( committed & terminal, torch.full_like(reasons, COMMIT_REASON_TERMINAL), reasons, ) return reasons 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 forward(self, hidden: torch.Tensor, seen: torch.Tensor, offsets: torch.Tensor) -> torch.Tensor: allowed = seen[:, None, :].expand(-1, hidden.shape[1], -1) if not self.global_attention: allowed = allowed & ((offsets[:, :, None] - offsets[:, None, :]).abs() < self.window) additive = torch.zeros(allowed.shape, device=hidden.device, dtype=torch.float32) additive = additive.masked_fill(~allowed, -1e9) normed = self.norm(hidden) hidden = hidden + self.attn(normed, normed, normed, attn_mask=additive) hidden = hidden + self.ff(self.ff_norm(hidden)) return torch.where(seen[..., None], hidden, torch.zeros_like(hidden)) class _WorkingBus(DecoderMemoryBus): def prepare_kv(self, memory: tuple[torch.Tensor, torch.Tensor], slot_identity: torch.Tensor | 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: torch.Tensor, reader_index: int, keys: torch.Tensor, values: torch.Tensor, seen: torch.Tensor, positions: torch.Tensor | None = None) -> torch.Tensor: written = super().read(hidden, reader_index, keys, values, positions, key_seen=seen) return torch.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 = nn.Parameter(torch.zeros(max(readers, 1), 1)) def reset_identity_parameters(self) -> None: with torch.no_grad(): self.alpha.zero_() def prepare_kv(self, memory: torch.Tensor, seen: torch.Tensor | None = None): if self.num_readers <= 0: return None return memory, seen def read(self, hidden: torch.Tensor, reader_index: int, memory: torch.Tensor, seen: torch.Tensor | None) -> torch.Tensor: delta = self._memory_ref().read_shared(memory, hidden) if seen is not None: delta = delta * seen[..., None].to(delta.dtype) return hidden + torch.tanh(self.alpha[reader_index]).to(hidden.dtype) * delta class LatentDeliberationTransformer(nn.Module): def __init__(self, *, hidden_size: int, latent_dim: int = 2816, ffn_dim: int = 7168, num_layers: int = 4, num_heads: int = 16, local_attention_window: int = 128, tape_probes: int = 4, history_kv_rank: int = 1024, num_memory_readers: int = 0, num_working_readers: int | None = None, num_persistent_readers: int | None = None, working_last_block_global: bool = True, 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) if self.packet_dim % num_heads: raise ValueError("Commit packet rank must be divisible by attention heads.") self.trajectory = GDN2TrajectoryMemory( hidden_size, probes=tape_probes, persistent_observation_dim=self.packet_dim, ) self.blocks = nn.ModuleList([ _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) readers = num_memory_readers if num_working_readers is None else num_working_readers persistent_readers = ( num_memory_readers if num_persistent_readers is None else num_persistent_readers ) bus_rank = history_kv_rank if history_kv_rank % num_heads == 0 else num_heads self.working_memory_bus = _WorkingBus( hidden_size, num_heads, readers, hidden_size, bus_rank, relative_bias=True, max_relative_span=max_canvas_length, ) self.persistent_memory_bus = _PersistentBus( self.trajectory.persistent, 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 reset_identity_parameters(self) -> None: self.working_memory_bus.reset_identity_parameters() self.persistent_memory_bus.reset_identity_parameters() def forward(self, *, token_embeddings: torch.Tensor, confidence: torch.Tensor, entropy: torch.Tensor, state: LatentDeliberationState, canvas_head: torch.Tensor | None = None) -> LatentProcessorOutput: 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 = torch.arange(canvas, device=token_embeddings.device)[None, :].expand(batch, -1) 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].to(hidden.dtype) next_state = replace(state, confidence=confidence.float(), entropy=entropy.float(), gdn2=memory) return LatentProcessorOutput(hidden, next_state) def observe_state(self, state: LatentDeliberationState, heavy: torch.Tensor, working: torch.Tensor, live: torch.Tensor, head: torch.Tensor) -> LatentDeliberationState: source = heavy.detach() + working updated = self.trajectory.observe(state.gdn2, source, live, head) return replace(state, gdn2=updated) def commit_write(self, *, memory: torch.Tensor, working_state: torch.Tensor, heavy_hidden: torch.Tensor, committed_token_embeddings: torch.Tensor, commit_lengths: torch.Tensor, commit_reason: torch.Tensor | None = None, canvas_head: torch.Tensor | None = None): batch, canvas, width = working_state.shape count = int(commit_lengths.max().item()) if count <= 0: return memory index = torch.arange(count, device=working_state.device)[None, :].expand(batch, -1) if canvas_head is not None: index = (index + canvas_head[:, None]) % canvas selected_working = working_state.gather( 1, index[..., None].expand(-1, -1, width) ) selected_heavy = heavy_hidden.detach().gather( 1, index[..., None].expand(-1, -1, width) ) if committed_token_embeddings.shape != selected_working.shape: raise ValueError("Committed embeddings do not match the prefix.") # Canvas processing already supplies bidirectional spatial context. # Separate normalized role channels feed the ordered GDN2 writer directly. # Unit-floor normalization keeps a zero Working state at zero without # amplifying its derivative by 1/sqrt(eps) on the first denoise. roles = tuple(F.rms_norm(value.float(), (width,), eps=1.0).to(value.dtype) for value in (selected_heavy, selected_working, committed_token_embeddings.detach())) packet = self.experience_in(torch.cat(roles, dim=-1)) reason = torch.zeros(batch, device=packet.device, dtype=torch.long) if commit_reason is None else commit_reason.long() reason = torch.where((reason == 1) | (reason == 2), 5, reason).clamp(0, 5) packet = packet + self.reason_embed(reason)[:, None] valid = torch.arange(count, device=packet.device)[None, :] < commit_lengths[:, None] written = self.trajectory.persistent.write_sequence(memory, packet, valid) return written