Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
block_wise_ssm.py47 linesDownload Raw Back to memory
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