TTVidT / tt3d.py
KBlueLeaf's picture
TT-VidT TT3D encoder (structured resample), transformers remote code
742c169 verified
Raw History Blame Contribute Delete
19.4 kB
"""
TemporalTransfer3D: Temporal attention with downsampled spatial context and 3D RoPE.
Each frame contributes M motion tokens + S downsampled spatial tokens.
Block-causal attention across time with 3D RoPE (x, y, t).
Both motion and spatial tokens receive attention residual.
Spatial tokens are downsampled via pixel unshuffle + a linear map,
and upsampled back via a linear map + pixel shuffle for residual add-back.
3D RoPE head_dim split: [x, y, t, unused] with 1/4 each.
Motion tokens use position (0, 0, t) β€” no spatial, only temporal.
Spatial tokens use position (x, y, t) β€” full 3D.
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .utils import compile_wrapper
from .layers import SwiGLU, GELUMLP, RMSNorm
# =============================================================================
# Block-causal mask (cached)
# =============================================================================
_mask_cache: dict[tuple, torch.Tensor] = {}
def get_block_causal_mask(t: int, tokens_per_frame: int, device: torch.device) -> torch.Tensor:
"""Bool mask [T*N, T*N] where frame i attends to frames 0..i."""
key = (t, tokens_per_frame, str(device))
if key not in _mask_cache:
causal = torch.tril(torch.ones(t, t, device=device, dtype=torch.bool))
n = tokens_per_frame
block = causal[:, :, None, None].expand(-1, -1, n, n)
mask = block.permute(0, 2, 1, 3).reshape(t * n, t * n)
_mask_cache[key] = mask
return _mask_cache[key]
# =============================================================================
# 3D RoPE
# =============================================================================
def _compute_freqs(dim: int, max_period: float = 10000.0) -> torch.Tensor:
"""Frequency bands for RoPE. Returns [dim//2]."""
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(half, dtype=torch.float32) / half
)
return freqs
def _apply_rope_1d(x: torch.Tensor, freqs: torch.Tensor, positions: torch.Tensor) -> torch.Tensor:
"""
Apply RoPE to a slice of x along the last dim.
Args:
x: [..., dim] β€” the slice of q or k to rotate
freqs: [dim//2] β€” precomputed frequency bands
positions: [...] β€” positions for each token
Returns:
[..., dim] rotated tensor
"""
half = x.shape[-1] // 2
angles = positions.unsqueeze(-1).float() * freqs.to(x.device)
cos = torch.cos(angles).to(x.dtype)
sin = torch.sin(angles).to(x.dtype)
x1, x2 = x[..., :half], x[..., half:]
return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1)
class RoPE3D(nn.Module):
"""
3D Rotary Position Embedding for (x, y, t) positions.
Head dim split into 4 equal parts: [x_rope, y_rope, t_rope, unused].
The unused portion passes through unchanged (identity).
"""
def __init__(self, head_dim: int, max_period: float = 10000.0):
super().__init__()
self.head_dim = head_dim
self.dim_x = head_dim // 4
self.dim_y = head_dim // 4
self.dim_t = head_dim // 4
self.dim_unused = head_dim - self.dim_x - self.dim_y - self.dim_t
self.max_period = max_period
self.register_buffer("freqs_x", _compute_freqs(self.dim_x, max_period), persistent=False)
self.register_buffer("freqs_y", _compute_freqs(self.dim_y, max_period), persistent=False)
self.register_buffer("freqs_t", _compute_freqs(self.dim_t, max_period), persistent=False)
def reset_buffers(self) -> None:
"""Recompute the non-persistent buffers (e.g. after transformers' meta-device loading)."""
self.freqs_x.copy_(_compute_freqs(self.dim_x, self.max_period))
self.freqs_y.copy_(_compute_freqs(self.dim_y, self.max_period))
self.freqs_t.copy_(_compute_freqs(self.dim_t, self.max_period))
def forward(self, q: torch.Tensor, k: torch.Tensor, positions: torch.Tensor):
"""
Args:
q, k: [B, H, S, head_dim]
positions: [S, 3] β€” (x, y, t) per token
Returns:
q_rot, k_rot: same shape
"""
pos_x = positions[:, 0]
pos_y = positions[:, 1]
pos_t = positions[:, 2]
def apply(tensor):
t_x = tensor[..., :self.dim_x]
t_y = tensor[..., self.dim_x:self.dim_x + self.dim_y]
t_t = tensor[..., self.dim_x + self.dim_y:self.dim_x + self.dim_y + self.dim_t]
t_u = tensor[..., self.dim_x + self.dim_y + self.dim_t:]
t_x = _apply_rope_1d(t_x, self.freqs_x, pos_x)
t_y = _apply_rope_1d(t_y, self.freqs_y, pos_y)
t_t = _apply_rope_1d(t_t, self.freqs_t, pos_t)
return torch.cat([t_x, t_y, t_t, t_u], dim=-1)
return apply(q), apply(k)
# =============================================================================
# Pixel unshuffle/shuffle spatial resampling
# =============================================================================
def _hadamard(n: int) -> torch.Tensor:
h = torch.ones(1, 1, dtype=torch.float64)
while h.shape[0] < n:
h = torch.cat([torch.cat([h, h], 1), torch.cat([h, -h], 1)], 0)
if h.shape[0] != n:
raise ValueError(f"Walsh-Hadamard needs a power-of-two size, got {n}")
return h
def _dct(n: int) -> torch.Tensor:
"""Orthonormal DCT-II ``[frequency, index]``."""
k = torch.arange(n, dtype=torch.float64)[:, None]
c = torch.arange(n, dtype=torch.float64)[None]
m = torch.cos(math.pi * (c + 0.5) * k / n) * math.sqrt(2 / n)
m[0] /= math.sqrt(2)
return m
def _chirp_signs(n: int, i: int) -> torch.Tensor:
"""A deterministic +-1 pattern per frequency ``i`` (quadratic chirp, never 0)."""
c = torch.arange(n, dtype=torch.float64)
s = torch.sign(torch.cos(math.pi * (c * c * (2 * i + 1) + 3 * i * c) / n + 0.25 * i))
s[s == 0] = 1.0
return s
def structured_weight(dim: int, factor: int) -> torch.Tensor:
"""The fixed structured down weight ``[D, D*f^2]`` (column index ``c*f^2 + p``,
the pixel-unshuffle channel order).
Over the f^2 positions of a patch, take orthonormal Walsh-Hadamard frequencies
``z_i = sum_p h_i[p] x[:, p]``; then ``y = (1/f) sum_i Q_i z_i`` with
``Q_i = C^T diag(chi_i) C`` (C the orthonormal DCT-II over channels, chi_i
deterministic chirp signs): one orthogonal D x D map per position frequency.
Depends only on ``D`` and ``f``: nothing is stored.
"""
f2 = factor * factor
h = _hadamard(f2) / factor # [i, p], orthonormal rows
c = _dct(dim)
q = torch.stack([c.T @ (_chirp_signs(dim, i)[:, None] * c) for i in range(f2)])
w = torch.einsum("ioc,ip->ocp", q, h) / factor # [o, c, p]
return w.reshape(dim, dim * f2).float()
def _to_unshuffled(x: torch.Tensor, h: int, w: int, f: int) -> torch.Tensor:
"""[B, T, H*W, D] -> [B, T, H'W', D*f*f] (index c*f*f + p)."""
B, T = x.shape[:2]
x = x.unflatten(2, (h, w)).permute(0, 1, 4, 2, 3).flatten(0, 1) # [BT, D, H, W]
x = F.pixel_unshuffle(x, f) # [BT, D*f*f, H', W']
return x.flatten(2).transpose(1, 2).unflatten(0, (B, T)) # [B, T, H'W', D*f*f]
def _from_unshuffled(x: torch.Tensor, h: int, w: int, f: int) -> torch.Tensor:
"""[B, T, H'W', D*f*f] -> [B, T, H*W, D] (inverse of ``_to_unshuffled``)."""
B, T = x.shape[:2]
x = x.transpose(-1, -2).unflatten(-1, (h // f, w // f)).flatten(0, 1) # [BT, D*f*f, H', W']
x = F.pixel_shuffle(x, f) # [BT, D, H, W]
return x.unflatten(0, (B, T)).flatten(3, 4).transpose(-1, -2) # [B, T, H*W, D]
class SpatialDownsample(nn.Module):
"""
[B, T, H*W, D] -> [B, T, (H/f)*(W/f), D]: pixel unshuffle, the fixed dense
``structured_weight`` (no parameters, rebuilt from D and f, not saved), then a
trainable D x D channel mix ``I + mix`` (``mix`` zero at init). Params: DΒ².
"""
def __init__(self, hidden_size: int, factor: int):
super().__init__()
self.hidden_size = hidden_size
self.factor = factor
self.register_buffer("fixed", self._fixed(), persistent=False)
self.mix = nn.Parameter(torch.zeros(hidden_size, hidden_size))
def _fixed(self) -> torch.Tensor:
return structured_weight(self.hidden_size, self.factor)
def reset_parameters(self) -> None:
"""Channel mix back to the identity (after a generic init such as mup_init)."""
nn.init.zeros_(self.mix)
def reset_buffers(self) -> None:
"""Recompute the fixed weight (e.g. after transformers' meta-device loading)."""
self.fixed.copy_(self._fixed())
def forward(self, x: torch.Tensor, h: int, w: int) -> torch.Tensor:
"""
Args:
x: [B, T, H*W, D]
h, w: spatial grid dims
Returns:
[B, T, (H//f)*(W//f), D]
"""
y = F.linear(_to_unshuffled(x, h, w, self.factor), self.fixed) # [B, T, H'W', D]
return y + F.linear(y, self.mix)
class SpatialUpsample(nn.Module):
"""
[B, T, (H/f)*(W/f), D] -> [B, T, H*W, D], the counterpart of SpatialDownsample:
channel mix ``I + mix``, then ``f`` times the transpose of the fixed down weight
(writes go back in the basis that was read), pixel shuffle. Params: DΒ².
"""
def __init__(self, hidden_size: int, factor: int):
super().__init__()
self.hidden_size = hidden_size
self.factor = factor
self.register_buffer("fixed", self._fixed(), persistent=False)
self.mix = nn.Parameter(torch.zeros(hidden_size, hidden_size))
def _fixed(self) -> torch.Tensor:
return self.factor * structured_weight(self.hidden_size, self.factor).T.contiguous()
def reset_parameters(self) -> None:
"""Channel mix back to the identity (after a generic init such as mup_init)."""
nn.init.zeros_(self.mix)
def reset_buffers(self) -> None:
"""Recompute the fixed weight (e.g. after transformers' meta-device loading)."""
self.fixed.copy_(self._fixed())
def forward(self, x: torch.Tensor, h: int, w: int) -> torch.Tensor:
"""
Args:
x: [B, T, H'Β·W', D] where H' = h // f, W' = w // f
h, w: ORIGINAL spatial grid dims (output size)
Returns:
[B, T, H*W, D]
"""
x = x + F.linear(x, self.mix)
return _from_unshuffled(F.linear(x, self.fixed), h, w, self.factor)
# =============================================================================
# TemporalTransfer3D
# =============================================================================
class TemporalTransfer3D(nn.Module):
"""
Temporal attention with downsampled spatial context and 3D RoPE.
Both motion and spatial tokens participate in block-causal attention.
Spatial tokens are downsampled (pixel unshuffle + linear), attend temporally,
then upsampled (linear + pixel shuffle) for residual add-back to the spatial stream.
Args:
hidden_size: model dimension
intermediate_size: FFN hidden dim
num_heads: attention heads
downsample_factor: spatial downsample factor (e.g. 4 β†’ 16x16 β†’ 4x4)
ffn_type: "swiglu" or "gelu"
qk_norm: use cosine-similarity attention
"""
def __init__(
self,
hidden_size: int,
intermediate_size: int,
num_heads: int,
downsample_factor: int = 4,
ffn_type: str = "swiglu",
qk_norm: bool = False,
):
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
self.downsample_factor = downsample_factor
self.qk_norm = qk_norm
# Spatial resampling: fixed structured weight + D x D channel mix
self.spatial_down = SpatialDownsample(hidden_size, downsample_factor)
self.spatial_up = SpatialUpsample(hidden_size, downsample_factor)
# Zero-init output projection for spatial add-back (identity at init)
self.spatial_out_proj = nn.Linear(hidden_size, hidden_size, bias=False)
nn.init.zeros_(self.spatial_out_proj.weight)
# Pre-norm
self.norm1 = RMSNorm(hidden_size)
self.norm2 = RMSNorm(hidden_size)
# Attention projections
self.q_proj = nn.Linear(hidden_size, hidden_size)
self.k_proj = nn.Linear(hidden_size, hidden_size)
self.v_proj = nn.Linear(hidden_size, hidden_size)
self.out_proj = nn.Linear(hidden_size, hidden_size)
if qk_norm:
self.qk_scale = nn.Parameter(torch.full([num_heads, 1, 1], 10.0))
# 3D RoPE
self.rope = RoPE3D(self.head_dim)
# FFN
if ffn_type == "gelu":
self.mlp = GELUMLP(hidden_size, intermediate_size)
else:
self.mlp = SwiGLU(hidden_size, intermediate_size)
# Position cache
self._pos_cache: dict[tuple, torch.Tensor] = {}
def _build_positions(
self, M: int, ds_h: int, ds_w: int, T: int, device: torch.device
) -> torch.Tensor:
"""
Build 3D positions for all tokens across all frames.
Per-frame layout: [MT_0..MT_{M-1}, DS_(0,0)..DS_(ds_w-1,ds_h-1)]
Motion: (0, 0, t) Spatial: (x, y, t)
Returns: [T * (M + ds_h*ds_w), 3]
"""
key = (M, ds_h, ds_w, T, str(device))
if key in self._pos_cache:
return self._pos_cache[key]
tokens_per_frame = M + ds_h * ds_w
positions = torch.zeros(T, tokens_per_frame, 3, device=device)
for t in range(T):
# Motion tokens: (0, 0, t)
positions[t, :M, 2] = t
# Spatial tokens: (x, y, t)
idx = M
for row in range(ds_h):
for col in range(ds_w):
positions[t, idx, 0] = col
positions[t, idx, 1] = row
positions[t, idx, 2] = t
idx += 1
positions = positions.reshape(T * tokens_per_frame, 3)
self._pos_cache[key] = positions
return positions
@compile_wrapper
def _attention(self, x: torch.Tensor, positions: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
"""
Block-causal attention with 3D RoPE.
Args:
x: [B, S, D] β€” normed, flattened (T * tokens_per_frame)
positions: [S, 3] β€” (x, y, t) per token
mask: [S, S] β€” block-causal bool mask
"""
from .layers import _qk_norm
B, S, D = x.shape
q = self.q_proj(x).view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
# 3D RoPE
q, k = self.rope(q, k, positions)
# QK-norm (cosine-sim attention)
if self.qk_norm:
q, k = _qk_norm(q, k, self.qk_scale)
attn = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, scale=1.0)
else:
attn = F.scaled_dot_product_attention(q, k, v, attn_mask=mask)
return self.out_proj(attn.transpose(1, 2).reshape(B, S, D))
def forward(
self,
motion_tokens: torch.Tensor,
spatial_tokens: torch.Tensor,
spatial_h: int,
spatial_w: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Args:
motion_tokens: [B, T, M, D] β€” motion tokens from DINOv3 stream
spatial_tokens: [B, T, H*W, D] β€” spatial tokens from DINOv3 stream
spatial_h, spatial_w: spatial grid dimensions (e.g. 16, 16)
Returns:
motion_tokens: [B, T, M, D] β€” enriched motion tokens
spatial_tokens: [B, T, H*W, D] β€” spatial tokens with temporal residual
"""
B, T, M, D = motion_tokens.shape
f = self.downsample_factor
ds_h, ds_w = spatial_h // f, spatial_w // f
tokens_per_frame = M + ds_h * ds_w
# 1. Downsample spatial: pixel unshuffle + linear
ds_spatial = self.spatial_down(spatial_tokens, spatial_h, spatial_w)
# ds_spatial: [B, T, ds_h*ds_w, D]
# 2. Concat motion + downsampled spatial per frame
x = torch.cat([motion_tokens, ds_spatial], dim=2) # [B, T, M+S_down, D]
# 3. Pre-norm + flatten temporal dim
x_normed = self.norm1(x).flatten(1, 2) # [B, T*(M+S_down), D]
# 4. Positions and mask
positions = self._build_positions(M, ds_h, ds_w, T, x.device)
mask = get_block_causal_mask(T, tokens_per_frame, x.device)
# 5. Block-causal attention with 3D RoPE
attn_out = self._attention(x_normed, positions, mask)
attn_out = attn_out.unflatten(1, (T, tokens_per_frame))
# 6. Attention residual on ALL tokens (motion + spatial)
x = x + attn_out
# 7. FFN on ALL tokens (motion + spatial)
x = x + self.mlp(self.norm2(x))
# 8. Split back into motion and spatial
motion_tokens = x[:, :, :M]
spatial_part = x[:, :, M:]
# 9. Spatial add-back: zero-init proj β†’ upsample β†’ residual
# spatial_out_proj is zero-init, so initially this is a no-op
spatial_residual = self.spatial_up(
self.spatial_out_proj(spatial_part), spatial_h, spatial_w
)
spatial_tokens = spatial_tokens + spatial_residual
return motion_tokens, spatial_tokens
# =============================================================================
# Smoke test
# =============================================================================
if __name__ == "__main__":
B, T, M, D = 2, 8, 8, 768
H, W = 16, 16
layer = TemporalTransfer3D(
hidden_size=D,
intermediate_size=D * 4,
num_heads=12,
downsample_factor=4,
)
motion = torch.randn(B, T, M, D)
spatial = torch.randn(B, T, H * W, D)
motion_out, spatial_out = layer(motion, spatial, H, W)
print(f"Input: motion={motion.shape}, spatial={spatial.shape}")
print(f"Output: motion={motion_out.shape}, spatial={spatial_out.shape}")
assert motion_out.shape == (B, T, M, D)
assert spatial_out.shape == (B, T, H * W, D)
# Verify spatial add-back is zero-init (residual starts as no-op)
diff = (spatial_out - spatial).abs().max().item()
print(f"Spatial diff at init (should be ~0 from zero-init out proj): {diff:.6f}")
# Check gradient flows through spatial
motion.requires_grad_(True)
spatial.requires_grad_(True)
m_out, s_out = layer(motion, spatial, H, W)
loss = m_out.sum() + s_out.sum()
loss.backward()
print(f"Gradient on spatial: {spatial.grad is not None}, norm={spatial.grad.norm():.4f}")
print(f"Gradient on motion: {motion.grad is not None}, norm={motion.grad.norm():.4f}")
# Different downsample factors
for ds in [1, 2, 4, 8]:
l = TemporalTransfer3D(D, D * 4, 12, downsample_factor=ds)
m_o, s_o = l(motion.detach(), spatial.detach(), H, W)
ds_tokens = (H // ds) * (W // ds)
print(f" ds={ds}: {ds_tokens} spatial tokens/frame, total={M + ds_tokens}/frame")
print("\nSmoke test passed!")