Team Ai
Apppublic

hugging-apps/echo-memory

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