"""Reference GDN2 matrix memory for trajectory-time recurrence. The state is FP32. Leading dimensions are independent rows or canvas cells; only the last three dimensions, [heads, key, value], belong to the rule. """ from __future__ import annotations import torch from torch import nn from torch.nn import functional as F class _GateProjection(nn.Module): """GDN2 gates need channel control, not a second full-width content map.""" def __init__(self, source: int, target: int, rank: int) -> None: super().__init__() self.down = nn.Linear(source, rank, bias=False) self.up = nn.Linear(rank, target, bias=False) def forward(self, source: torch.Tensor) -> torch.Tensor: return self.up(self.down(source)) class GDN2Memory(nn.Module): def __init__(self, input_dim: int, heads: int, key_dim: int, value_dim: int, *, observation_dim: int | None = None) -> None: super().__init__() if min(input_dim, heads, key_dim, value_dim) <= 0: raise ValueError("GDN2 dimensions must be positive.") self.heads, self.key_dim, self.value_dim = heads, key_dim, value_dim observation_dim = input_dim if observation_dim is None else observation_dim if observation_dim <= 0: raise ValueError("GDN2 observation width must be positive.") self.q_proj = nn.Linear(input_dim, heads * key_dim, bias=False) self.k_proj = nn.Linear(observation_dim, heads * key_dim, bias=False) self.v_proj = nn.Linear(observation_dim, heads * value_dim, bias=False) self.f_proj = _GateProjection(observation_dim, heads * key_dim, min(observation_dim, key_dim)) self.b_proj = nn.Linear(observation_dim, heads * key_dim, bias=False) self.w_proj = nn.Linear(observation_dim, heads * value_dim, bias=False) self.g_proj = _GateProjection(input_dim, heads * value_dim, min(input_dim, value_dim)) self.o_proj = nn.Linear(heads * value_dim, input_dim, bias=False) self.a_log = nn.Parameter(torch.zeros(heads)) self.dt_bias = nn.Parameter(torch.full((heads, key_dim), -6.906255)) def _projections(self, source: torch.Tensor): source = F.rms_norm(source.float(), (source.shape[-1],), eps=1e-6).to(source.dtype) shape = source.shape[:-1] key_shape = (*shape, self.heads, self.key_dim) value_shape = (*shape, self.heads, self.value_dim) k = F.normalize(F.silu(self.k_proj(source).float()).reshape(key_shape), dim=-1) v = F.silu(self.v_proj(source).float()).reshape(value_shape) head_rate = self.a_log.float().exp().reshape(*((1,) * (source.ndim - 1)), self.heads, 1) decay = torch.exp(-head_rate * F.softplus( self.f_proj(source).float().reshape(key_shape) + self.dt_bias.float())) erase = torch.sigmoid(self.b_proj(source).float().reshape(key_shape)) write = torch.sigmoid(self.w_proj(source).float().reshape(value_shape)) return k, v, decay, erase, write def _query(self, source: torch.Tensor) -> torch.Tensor: return F.normalize(F.silu(self.q_proj(source).float()).reshape( *source.shape[:-1], self.heads, self.key_dim), dim=-1) def _output(self, value: torch.Tensor, source: torch.Tensor) -> torch.Tensor: gate = F.silu(self.g_proj(source).float().reshape(value.shape)) value = F.rms_norm(value, (self.value_dim,), eps=1e-6) * gate return self.o_proj(value.flatten(-2).to(source.dtype)) def read(self, state: torch.Tensor, source: torch.Tensor) -> torch.Tensor: if state.shape != (*source.shape[:-1], self.heads, self.key_dim, self.value_dim): raise ValueError("GDN2 state and query leading dimensions differ.") value = torch.einsum("...hk,...hkv->...hv", self._query(source), state.float()) return self._output(value, source) def read_shared(self, state: torch.Tensor, source: torch.Tensor) -> torch.Tensor: """Read one row matrix for all queries without a canvas-sized matrix product.""" if source.ndim != 3 or state.shape != (source.shape[0], self.heads, self.key_dim, self.value_dim): raise ValueError("Shared GDN2 read requires [batch, canvas, width] queries.") query = self._query(source).transpose(1, 2) value = (query @ state.float()).transpose(1, 2) return self._output(value, source) @staticmethod def _transition(state, k, v, decay, erase, write, valid): decayed = state.float() * decay.unsqueeze(-1) old = torch.einsum("...hk,...hkv->...hv", erase * k, decayed) candidate = decayed + k.unsqueeze(-1) * (write * v - old).unsqueeze(-2) return candidate if valid is None else torch.where( valid[..., None, None, None], candidate, state.float()) def transition(self, state: torch.Tensor, source: torch.Tensor, valid: torch.Tensor | None = None) -> torch.Tensor: if state.shape != (*source.shape[:-1], self.heads, self.key_dim, self.value_dim): raise ValueError("GDN2 state and observation leading dimensions differ.") k, v, decay, erase, write = self._projections(source) if valid is not None: if valid.shape != source.shape[:-1]: raise ValueError("GDN2 valid mask must match observation rows.") return self._transition(state, k, v, decay, erase, write, valid) def write_sequence(self, state: torch.Tensor, source: torch.Tensor, valid: torch.Tensor) -> torch.Tensor: """Project a commit packet once, then preserve ordered GDN2 recurrence.""" if source.ndim != 3 or valid.shape != source.shape[:2]: raise ValueError("GDN2 sequence and mask must share [batch, length].") if state.shape != (source.shape[0], self.heads, self.key_dim, self.value_dim): raise ValueError("GDN2 sequence state shape differs.") projected = self._projections(source) for index in range(source.shape[1]): state = self._transition(state, *(part[:, index] for part in projected), valid[:, index]) return state