"""Dual-timescale matrix state with denoise and commit lifetimes. This module is independent of the decoder and commit policy. It provides the reference state transition used when replacing the old history and slot paths. """ from __future__ import annotations from dataclasses import dataclass import torch from torch import nn from .gdn2_memory import GDN2Memory @dataclass(frozen=True) class GDN2TrajectoryState: cells: torch.Tensor row: torch.Tensor persistent: torch.Tensor seen: torch.Tensor def shift(self, lengths: torch.Tensor) -> "GDN2TrajectoryState": """Shift a non-ring canvas and zero its newly filled tail.""" batch, canvas = self.seen.shape if lengths.shape != (batch,): raise ValueError("Commit lengths must be per row.") physical = torch.arange(canvas, device=self.seen.device)[None, :] + lengths[:, None] kept = physical < canvas selected = physical.clamp_max(canvas - 1) cells = self.cells.gather( 1, selected[..., None, None, None].expand_as(self.cells) ) seen = self.seen.gather(1, selected) return GDN2TrajectoryState( cells.masked_fill(~kept[..., None, None, None], 0.0), self.row, self.persistent, seen & kept, ) class GDN2TrajectoryMemory(nn.Module): def __init__(self, width: int, *, probes: int = 4, working_heads: int = 16, working_key: int = 64, working_value: int = 64, persistent_heads: int = 16, persistent_key: int = 128, persistent_value: int = 128, persistent_observation_dim: int | None = None) -> None: super().__init__() if probes <= 0: raise ValueError("Probe count must be positive.") self.probes = probes self.cell = GDN2Memory(width, working_heads, working_key, working_value) self.row = GDN2Memory(width, working_heads, working_key, working_value) self.persistent = GDN2Memory(width, persistent_heads, persistent_key, persistent_value, observation_dim=persistent_observation_dim) self.probe_embed = nn.Embedding(probes, width) nn.init.normal_(self.probe_embed.weight, std=0.02) def read(self, state: GDN2TrajectoryState, query: torch.Tensor) -> torch.Tensor: result = (self.cell.read(state.cells, query) + self.row.read_shared(state.row, query) + self.persistent.read_shared(state.persistent, query)) return torch.where(state.seen[..., None], result, torch.zeros_like(result)) def observe(self, state: GDN2TrajectoryState, observation: torch.Tensor, live: torch.Tensor, head: torch.Tensor) -> GDN2TrajectoryState: batch, canvas, width = observation.shape if live.shape != (batch, canvas) or head.shape != (batch,): raise ValueError("Observation mask and head have incorrect shapes.") cells = self.cell.transition(state.cells, observation, live) seen = state.seen | live logical_idx = (head[:, None] + torch.arange(canvas, device=head.device)[None, :]) % canvas logical = observation.gather(1, logical_idx[..., None].expand(-1, -1, width)) logical_live = live.gather(1, logical_idx) row_state = state.row for probe in range(self.probes): lo = canvas * probe // self.probes hi = canvas * (probe + 1) // self.probes selected = logical_live[:, lo:hi] count = selected.sum(dim=1, keepdim=True) pooled = (logical[:, lo:hi].float() * selected[..., None]).sum(dim=1) pooled = (pooled / count.clamp_min(1)).to(observation.dtype) pooled = pooled + self.probe_embed.weight[probe].to(observation.dtype) row_state = self.row.transition(row_state, pooled, count[:, 0] > 0) return GDN2TrajectoryState(cells, row_state, state.persistent, seen)