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