hugging-apps/echo-memory
0
1import torch2import torch.nn.functional as F3 4 5def _compress_weights(ratio: int, strategy: str = "distance_merge", recent_keep_ratio: float = 0.5, device=None, dtype=None):6 if ratio <= 1:7 return None8 strategy = str(strategy or "distance_merge").lower()9 # Baseline-aligned default: non-overlapping mean pool on each r-frame group.10 if strategy in ("distance_merge", "mean", "uniform"):11 return None12 if strategy in ("recent_weighted", "weighted_recent"):13 # Optional weighted variant (kept for compatibility experiments).14 idx = torch.arange(ratio, device=device, dtype=dtype)15 w = (1.0 - float(recent_keep_ratio)) + float(recent_keep_ratio) * ((idx + 1.0) / float(ratio))16 w = w / w.sum()17 return w18 return torch.full((ratio,), 1.0 / float(ratio), device=device, dtype=dtype)19 20 21def framepack_length_compress_context_latents(22 context_latents: torch.Tensor,23 framepack_ratio: int,24 strategy: str = "distance_merge",25 recent_keep_ratio: float = 0.5,26 multiscale_w2: float = 0.25,27 multiscale_w4: float = 0.15,28):29 # context_latents: (B, C, K, H, W)30 if context_latents is None:31 return None, 0, 0, 032 if context_latents.ndim != 5:33 raise ValueError(f"context_latents must be 5D (B,C,K,H,W), got {tuple(context_latents.shape)}")34 r = int(framepack_ratio)35 if r <= 1:36 k = int(context_latents.shape[2])37 return context_latents, k, k, k38 39 b, c, k_orig, h, w = context_latents.shape40 pad = (r - (k_orig % r)) % r41 if pad > 0:42 pad_lat = context_latents[:, :, -1:, :, :].repeat(1, 1, pad, 1, 1)43 context_latents = torch.cat([context_latents, pad_lat], dim=2)44 k_pad = int(context_latents.shape[2])45 new_k = k_pad // r46 47 grouped = context_latents.reshape(b, c, new_k, r, h, w)48 strategy = str(strategy or "distance_merge").lower()49 if strategy in ("packed_multiscale", "multiscale_packed", "multi_scale_packed"):50 base = grouped.mean(dim=3)51 52 # Base-code inspired approximation: aggregate history with extra low-res spatial views53 # (1x/2x/4x) and fuse back to the packed latent stream.54 x2 = F.avg_pool3d(context_latents, kernel_size=(1, 2, 2), stride=(1, 2, 2))55 x4 = F.avg_pool3d(context_latents, kernel_size=(1, 4, 4), stride=(1, 4, 4))56 x2 = F.interpolate(x2, size=(k_pad, h, w), mode="trilinear", align_corners=False)57 x4 = F.interpolate(x4, size=(k_pad, h, w), mode="trilinear", align_corners=False)58 b2 = x2.reshape(b, c, new_k, r, h, w).mean(dim=3)59 b4 = x4.reshape(b, c, new_k, r, h, w).mean(dim=3)60 w2 = float(multiscale_w2 or 0.0)61 w4 = float(multiscale_w4 or 0.0)62 w1 = max(1e-6, 1.0 - w2 - w4)63 s = w1 + w2 + w464 out = (w1 * base + w2 * b2 + w4 * b4) / s65 else:66 cw = _compress_weights(r, strategy=strategy, recent_keep_ratio=recent_keep_ratio, device=context_latents.device, dtype=context_latents.dtype)67 if cw is None:68 out = grouped.mean(dim=3)69 else:70 out = (grouped * cw.view(1, 1, 1, r, 1, 1)).sum(dim=3)71 return out, int(new_k), int(k_pad), int(k_orig)72 73 74def framepack_align_context_actions_to_latents(75 context_actions,76 K_orig_latent: int,77 K_after_pad: int,78 framepack_ratio: int,79 device=None,80 dtype=None,81 strategy: str = "distance_merge",82 recent_keep_ratio: float = 0.5,83):84 if context_actions is None:85 return None86 x = context_actions87 if not isinstance(x, torch.Tensor):88 x = torch.tensor(x, device=device, dtype=dtype or torch.float32)89 else:90 if device is not None:91 x = x.to(device=device)92 if dtype is not None:93 x = x.to(dtype=dtype)94 if x.ndim not in (2, 3):95 raise ValueError(f"context_actions must be 2D/3D, got shape {tuple(x.shape)}")96 r = int(framepack_ratio)97 if r <= 1:98 return x99 100 if x.ndim == 2:101 # (K, D)102 k, d = x.shape103 k_expected = int(K_orig_latent)104 if k < k_expected:105 raise ValueError(f"context_actions shorter than K_orig_latent: {k} < {k_expected}")106 x = x[:k_expected, :]107 pad = int(K_after_pad) - k_expected108 if pad > 0:109 x = torch.cat([x, x[-1:, :].repeat(pad, 1)], dim=0)110 new_k = int(K_after_pad) // r111 grouped = x.reshape(new_k, r, d)112 cw = _compress_weights(r, strategy=str(strategy or "distance_merge").lower(), recent_keep_ratio=recent_keep_ratio, device=x.device, dtype=x.dtype)113 return grouped.mean(dim=1) if cw is None else (grouped * cw.view(1, r, 1)).sum(dim=1)114 115 # (B, K, D)116 b, k, d = x.shape117 k_expected = int(K_orig_latent)118 if k < k_expected:119 raise ValueError(f"context_actions shorter than K_orig_latent: {k} < {k_expected}")120 x = x[:, :k_expected, :]121 pad = int(K_after_pad) - k_expected122 if pad > 0:123 x = torch.cat([x, x[:, -1:, :].repeat(1, pad, 1)], dim=1)124 new_k = int(K_after_pad) // r125 grouped = x.reshape(b, new_k, r, d)126 cw = _compress_weights(r, strategy=str(strategy or "distance_merge").lower(), recent_keep_ratio=recent_keep_ratio, device=x.device, dtype=x.dtype)127 return grouped.mean(dim=2) if cw is None else (grouped * cw.view(1, 1, r, 1)).sum(dim=2)128 129 