Modilify-Mk2-preview / gdn2_memory.py
ydy9038074's picture
Publish Modilify Mk2 Preview step 1250 schema25
b88f761 verified
Raw History Blame Contribute Delete
6.17 kB
"""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