prism / encoder.py
litcoderr's picture
Publish PRISM weights and modeling code
a596b0a verified
Raw
History Blame Contribute Delete
5.57 kB
"""Decompositional Encoder θ.
Decomposes a video into a view-invariant stream ``z_vi`` and a view-variant
stream ``z_vv`` (paper §3.1). Two sub-modules:
- ``QFormer`` : per-frame BLIP-2-style Q-Former with two learned
queries (Q_vi, Q_vv) that attend to a single
frame's frozen patch tokens.
- ``CausalTemporalEncoder`` : two parallel causal temporal streams (one per
factor). Within a stream, frame ``t`` attends to
frames ``≤ t``; the streams never attend to each
other (cross-factor mixing is deferred to φ).
``forward(patches) -> (z_vi, z_vv)`` with each ``(B, T, d_z)``.
"""
from __future__ import annotations
import torch
from torch import nn
from .layers import QFormerBlock, TemporalBlock, build_sin_pos_embed, causal_mask
class QFormer(nn.Module):
"""Per-frame Q-Former with N=2 queries: query 0 → z_vi, query 1 → z_vv."""
def __init__(
self,
d_z: int = 512,
d_kv: int = 1024,
depth: int = 4,
num_heads: int = 8,
mlp_ratio: float = 4.0,
):
super().__init__()
self.d_z = d_z
self.d_kv = d_kv
self.queries = nn.Parameter(torch.zeros(1, 2, d_z))
nn.init.normal_(self.queries, std=0.02)
self.blocks = nn.ModuleList(
[QFormerBlock(d_z, d_kv, num_heads, mlp_ratio) for _ in range(depth)]
)
self.final_norm = nn.LayerNorm(d_z)
def forward(self, patches: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""patches: ``(B, T, P, d_kv)`` → ``(z_vi, z_vv)`` each ``(B, T, d_z)``."""
B, T, P, d_kv = patches.shape
assert d_kv == self.d_kv, f"expected d_kv={self.d_kv}, got {d_kv}"
kv = patches.reshape(B * T, P, d_kv)
q = self.queries.expand(B * T, -1, -1).contiguous()
for block in self.blocks:
q = block(q, kv)
q = self.final_norm(q)
z_vi = q[:, 0, :].reshape(B, T, self.d_z) # query 0 → view-invariant
z_vv = q[:, 1, :].reshape(B, T, self.d_z) # query 1 → view-variant
return z_vi, z_vv
class CausalTemporalEncoder(nn.Module):
"""Two parallel causal temporal streams over z_vi and z_vv.
Each stream has its own (non-shared) stack of ``TemporalBlock``s. A shared
sinusoidal positional embedding is added before the blocks; a boolean
``key_padding_mask`` ``(B, T)`` blocks padded frames in both streams.
"""
def __init__(
self,
d_z: int = 512,
max_frames: int = 128,
depth: int = 12,
num_heads: int = 8,
mlp_ratio: float = 4.0,
):
super().__init__()
self.d_z = d_z
self.max_frames = max_frames
self.register_buffer(
"pos_embed", build_sin_pos_embed(max_frames, d_z), persistent=False
)
self.blocks_vi = nn.ModuleList(
[TemporalBlock(d_z, num_heads, mlp_ratio) for _ in range(depth)]
)
self.blocks_vv = nn.ModuleList(
[TemporalBlock(d_z, num_heads, mlp_ratio) for _ in range(depth)]
)
self.norm_vi = nn.LayerNorm(d_z)
self.norm_vv = nn.LayerNorm(d_z)
def _pos_embed(self, T: int, device, dtype) -> torch.Tensor:
# Training never exceeds max_frames; T > max_frames only happens when
# encoding a full native-fps video longer than the cap, where we extend
# the (deterministic) sinusoidal PE on the fly.
if T <= self.max_frames:
return self.pos_embed[:, :T, :]
return build_sin_pos_embed(T, self.d_z).to(device=device, dtype=dtype)
def forward(
self,
z_vi_seq: torch.Tensor,
z_vv_seq: torch.Tensor,
key_padding_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""z_vi_seq, z_vv_seq: ``(B, T, d_z)``; ``key_padding_mask`` True = padded."""
B, T, _ = z_vi_seq.shape
pos = self._pos_embed(T, z_vi_seq.device, z_vi_seq.dtype)
zi = z_vi_seq + pos
zv = z_vv_seq + pos
attn_mask = causal_mask(T, device=z_vi_seq.device)
for blk in self.blocks_vi:
zi = blk(zi, attn_mask=attn_mask, key_padding_mask=key_padding_mask)
for blk in self.blocks_vv:
zv = blk(zv, attn_mask=attn_mask, key_padding_mask=key_padding_mask)
return self.norm_vi(zi), self.norm_vv(zv)
class DecompositionalEncoder(nn.Module):
"""θ: video patches → (z_vi, z_vv).
Composes the per-frame ``QFormer`` with the ``CausalTemporalEncoder``.
"""
def __init__(
self,
d_z: int = 512,
d_kv: int = 1024,
qformer_depth: int = 4,
temporal_depth: int = 12,
num_heads: int = 8,
mlp_ratio: float = 4.0,
max_frames: int = 128,
):
super().__init__()
self.qformer = QFormer(
d_z=d_z, d_kv=d_kv, depth=qformer_depth,
num_heads=num_heads, mlp_ratio=mlp_ratio,
)
self.temporal = CausalTemporalEncoder(
d_z=d_z, max_frames=max_frames, depth=temporal_depth,
num_heads=num_heads, mlp_ratio=mlp_ratio,
)
def forward(
self,
patches: torch.Tensor,
key_padding_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
z_vi_pre, z_vv_pre = self.qformer(patches)
return self.temporal(z_vi_pre, z_vv_pre, key_padding_mask=key_padding_mask)