hugging-apps/echo-memory
0
1import torch2import torch.nn as nn3 4 5class BlockWiseStateSpaceMemory(nn.Module):6 """7 Paper-aligned block-wise recurrent SSM.8 9 This module is intentionally separate from VideoSSM hybrid. It performs a10 recurrent state update along the latent time axis for each spatial token11 trajectory, and is attached to selected DiT blocks.12 """13 14 def __init__(self, dim: int):15 super().__init__()16 self.dim = int(dim)17 self.in_proj = nn.Linear(self.dim, self.dim * 2)18 self.out_proj = nn.Linear(self.dim, self.dim)19 self.decay_logit = nn.Parameter(torch.zeros(self.dim))20 self.gate = nn.Parameter(torch.zeros(1))21 22 def forward(self, x: torch.Tensor, f: int, **_kwargs):23 # x: (B, F*S, D), where S is spatial tokens per latent frame.24 if x is None or x.ndim != 3:25 return x26 b, n, d = x.shape27 f = int(f or 0)28 if d != self.dim or f <= 1 or n % f != 0:29 return x30 31 spatial = n // f32 x_seq = x.reshape(b, f, spatial, d).permute(0, 2, 1, 3).reshape(b * spatial, f, d)33 update, update_gate = self.in_proj(x_seq).chunk(2, dim=-1)34 update = torch.tanh(update)35 update_gate = torch.sigmoid(update_gate)36 decay = torch.sigmoid(self.decay_logit).to(dtype=x.dtype, device=x.device).view(1, d)37 38 state = torch.zeros(x_seq.shape[0], d, dtype=x.dtype, device=x.device)39 outputs = []40 for t in range(f):41 state = decay * state + (1.0 - decay) * update[:, t, :]42 outputs.append(state * update_gate[:, t, :])43 y = torch.stack(outputs, dim=1)44 y = self.out_proj(y)45 y = y.reshape(b, spatial, f, d).permute(0, 2, 1, 3).reshape(b, n, d)46 return x + torch.tanh(self.gate) * y47 