Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
spatial_grid_memory.py87 linesDownload Raw Back to memory
1import torch2import torch.nn as nn3import torch.nn.functional as F4 5 6class SpatialGridMemory(nn.Module):7    def __init__(self, dim: int, grid_size: int = 8, num_tokens: int = 64):8        super().__init__()9        self.dim = int(dim)10        self.grid_size = int(grid_size)11        self.num_tokens = int(num_tokens)12        g2 = self.grid_size * self.grid_size13        # Keep key name aligned with ckpt loading in loop_utils.py (spatial_to_tokens).14        self.spatial_to_tokens = nn.Parameter(torch.zeros(g2, self.num_tokens))15        nn.init.normal_(self.spatial_to_tokens, std=0.02)16 17    @property18    def mix(self):19        # Backward compatibility for code that referenced the old attribute name.20        return self.spatial_to_tokens21 22    def forward(self, x_context: torch.Tensor, num_context_frames: int, h: int, w: int):23        # x_context: (B, K*H*W, D)24        if x_context is None or x_context.ndim != 3:25            return x_context26        b, n, d = x_context.shape27        if d != self.dim:28            raise ValueError(f"SpatialGridMemory dim mismatch: x={d} module={self.dim}")29        k = max(int(num_context_frames), 1)30        spatial = int(h) * int(w)31        if n != k * spatial:32            # Best effort fallback: treat x as a flat token map and pool directly.33            x_mean = x_context34        else:35            x_mean = x_context.reshape(b, k, spatial, d).mean(dim=1)  # (B, S, D)36 37        g2 = self.grid_size * self.grid_size38        pooled = F.adaptive_avg_pool1d(x_mean.transpose(1, 2), g2).transpose(1, 2)  # (B, G2, D)39        mix = torch.softmax(self.spatial_to_tokens, dim=0)  # (G2, M)40        mem = torch.einsum("bgd,gm->bmd", pooled, mix)  # (B, M, D)41        return mem42 43    def load_state_dict(self, state_dict, strict: bool = True):44        # Compatibility:45        # - old local key: mix46        # - current/baseline key: spatial_to_tokens47        sd = dict(state_dict)48        if "mix" in sd and "spatial_to_tokens" not in sd:49            sd["spatial_to_tokens"] = sd.pop("mix")50        # Ignore deprecated projection keys from prior experiments.51        sd.pop("out.weight", None)52        sd.pop("out.bias", None)53        return super().load_state_dict(sd, strict=False if not strict else strict)54 55 56class SpatialCrossAttnReadout(nn.Module):57    def __init__(self, dim: int, num_heads: int = 8):58        super().__init__()59        self.attn = nn.MultiheadAttention(embed_dim=int(dim), num_heads=int(num_heads), batch_first=True)60        self.gate = nn.Parameter(torch.zeros(1))61 62    def forward(self, x_target: torch.Tensor, mem_tokens: torch.Tensor):63        if x_target is None or mem_tokens is None:64            return x_target65        if x_target.numel() == 0 or mem_tokens.numel() == 0:66            return x_target67        delta, _ = self.attn(x_target, mem_tokens, mem_tokens, need_weights=False)68        return x_target + torch.tanh(self.gate) * delta69 70 71def apply_spatial_cross_attn_readout(x_target: torch.Tensor, mem_tokens: torch.Tensor, module: nn.Module = None):72    if module is None:73        module = SpatialCrossAttnReadout(dim=int(x_target.shape[-1]), num_heads=8).to(device=x_target.device, dtype=x_target.dtype)74    return module(x_target, mem_tokens)75 76 77def inject_spatial_memory(context: torch.Tensor, mem_tokens: torch.Tensor, mode: str = "concat_text"):78    mode = str(mode or "concat_text").lower()79    if mem_tokens is None or mode == "none":80        return context81    if context is None:82        return mem_tokens83    if mode in ("concat_text", "cross_attn_readout"):84        return torch.cat([context, mem_tokens], dim=1)85    return context86 87