TTVidT / pos_embed.py
KBlueLeaf's picture
TT-VidT TT3D encoder (structured resample), transformers remote code
742c169 verified
Raw History Blame Contribute Delete
12.6 kB
"""
Positional embedding modules for TT-VidT.
Includes:
- Sinusoidal 1D/2D position embeddings
- Adaptive RoPE (Rotary Position Embedding) with aspect-ratio-aware positions
- Temporal distance embeddings
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .utils import compile_wrapper
# ============================================================================
# Sinusoidal Position Embeddings
# ============================================================================
def get_1d_sincos_pos_embed(
embed_dim: int, pos: torch.Tensor, max_period: int = 10000
) -> torch.Tensor:
"""
Generate 1D sinusoidal positional embeddings.
Args:
embed_dim: Embedding dimension
pos: Position tensor [N]
max_period: Maximum period for frequency bands
Returns:
Sinusoidal embeddings [N, embed_dim]
"""
half_dim = embed_dim // 2
freqs = torch.exp(
-math.log(max_period)
* torch.arange(half_dim, dtype=torch.float32, device=pos.device)
/ half_dim
)
args = pos.unsqueeze(-1) * freqs # [N, D/2]
return torch.cat([torch.sin(args), torch.cos(args)], dim=-1) # [N, D]
def get_2d_sincos_pos_embed(
embed_dim: int, h: int, w: int, max_period: int = 10000
) -> torch.Tensor:
"""
Generate 2D sinusoidal positional embeddings.
Args:
embed_dim: Embedding dimension
h, w: Grid height and width
max_period: Maximum period for frequency bands
Returns:
Sinusoidal embeddings [H*W, embed_dim]
"""
grid_h = torch.arange(h, dtype=torch.float32)
grid_w = torch.arange(w, dtype=torch.float32)
grid = torch.stack(
torch.meshgrid(grid_h, grid_w, indexing="ij"), dim=-1
) # [H, W, 2]
grid = grid.reshape(-1, 2) # [H*W, 2]
# Split embedding dim for h and w
half_dim = embed_dim // 2
emb_h = get_1d_sincos_pos_embed(half_dim, grid[:, 0], max_period)
emb_w = get_1d_sincos_pos_embed(half_dim, grid[:, 1], max_period)
return torch.cat([emb_h, emb_w], dim=-1) # [H*W, D]
def get_adaptive_2d_pos(
h: int, w: int, device: torch.device
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Generate adaptive 2D positions with aspect-ratio-aware normalization.
Position mapping is normalized such that:
- max(H) * max(W) = 1 (positions are in normalized range)
- max(H) / max(W) = h / w (preserves aspect ratio)
This means: max_H = sqrt(h/w), max_W = sqrt(w/h)
Args:
h, w: Grid dimensions
device: Torch device
Returns:
(grid_h, grid_w): Flattened position grids, each [H*W]
"""
aspect = h / w
max_H = math.sqrt(aspect)
max_W = math.sqrt(1.0 / aspect)
pos_h = torch.linspace(0, max_H, h, device=device, dtype=torch.float32)
pos_w = torch.linspace(0, max_W, w, device=device, dtype=torch.float32)
grid_h, grid_w = torch.meshgrid(pos_h, pos_w, indexing="ij")
return grid_h.reshape(-1), grid_w.reshape(-1)
# ============================================================================
# Temporal Distance Embedding
# ============================================================================
class TemporalDistanceEmbedding(nn.Module):
"""
Temporal distance embedding as a concatenated token using sincos encoding.
Instead of adding temporal pos encoding, we create a dedicated token
that encodes the distance from source frame to target motion frame.
This token is concatenated: [frame_emb, motion_emb, time_info]
Uses sinusoidal positional encoding (like standard transformer timestep)
which generalizes to arbitrary frame distances without a fixed max.
"""
def __init__(self, hidden_size: int, max_period: int = 10000):
super().__init__()
self.hidden_size = hidden_size
self.max_period = max_period
# Precompute frequency bands for sincos encoding
half_dim = hidden_size // 2
freqs = torch.exp(
-math.log(max_period)
* torch.arange(half_dim, dtype=torch.float32)
/ half_dim
)
self.register_buffer("freqs", freqs)
def forward(
self,
batch_size: int,
num_frames: int,
device: torch.device,
start_distance: int = 1,
) -> torch.Tensor:
"""
Generate temporal distance tokens using sincos encoding.
Args:
batch_size: batch size
num_frames: number of motion frames (T)
device: torch device
start_distance: distance of first motion frame from reference
Returns:
[B, T, 1, D] temporal distance tokens
"""
distances = torch.arange(
start_distance,
start_distance + num_frames,
device=device,
dtype=torch.float32,
)
# Sincos encoding: [T, D/2] -> [T, D]
args = distances.unsqueeze(-1) * self.freqs.to(device) # [T, D/2]
time_tokens = torch.cat([torch.sin(args), torch.cos(args)], dim=-1) # [T, D]
# Expand: [1, T, 1, D] -> [B, T, 1, D]
return time_tokens.unsqueeze(0).unsqueeze(2).expand(batch_size, -1, -1, -1)
# ============================================================================
# Additive RoPE
# ============================================================================
class AdditiveRoPE2D(nn.Module):
"""
Additive 2D Rotary Position Embedding with adaptive aspect-ratio-aware positions.
x = x + rope(learnable_token, pos)
Position mapping is normalized such that:
- max(H) * max(W) = 1 (positions are in [0, 1] range)
- max(H) / max(W) = vid_h / vid_w (preserves aspect ratio)
This removes the need for max_h/max_w parameters and generalizes to any resolution.
"""
def __init__(self, hidden_size: int, max_period: int = 10000):
super().__init__()
self.hidden_size = hidden_size
self.max_period = max_period
# Learnable token that gets modulated by position
self.learnable_token = nn.Parameter(torch.randn(1, 1, hidden_size) * 0.02)
# Precompute frequency bands for sincos encoding (half for h, half for w)
quarter_dim = hidden_size // 4
freqs = torch.exp(
-math.log(max_period)
* torch.arange(quarter_dim, dtype=torch.float32)
/ quarter_dim
)
self.register_buffer("freqs", freqs)
def _get_adaptive_pos_embed(
self, h: int, w: int, device: torch.device
) -> torch.Tensor:
"""Generate position embeddings with adaptive aspect-ratio-aware mapping."""
grid_h, grid_w = get_adaptive_2d_pos(h, w, device)
# Sincos encoding for h and w positions
args_h = grid_h.unsqueeze(-1) * self.freqs.to(device) # [L, D/4]
emb_h = torch.cat([torch.sin(args_h), torch.cos(args_h)], dim=-1) # [L, D/2]
args_w = grid_w.unsqueeze(-1) * self.freqs.to(device) # [L, D/4]
emb_w = torch.cat([torch.sin(args_w), torch.cos(args_w)], dim=-1) # [L, D/2]
return torch.cat([emb_h, emb_w], dim=-1) # [L, D]
def forward(self, x: torch.Tensor, h: int, w: int) -> torch.Tensor:
"""
Args:
x: [B, T, L, D] or [B, L, D] input tensor
h, w: spatial dimensions
Returns:
x with additive RoPE applied
"""
pos = self._get_adaptive_pos_embed(h, w, x.device).unsqueeze(0) # [1, L, D]
rope_emb = self.learnable_token * pos # [1, L, D]
if x.ndim == 4:
rope_emb = rope_emb.unsqueeze(1) # [1, 1, L, D]
return x + rope_emb
# ============================================================================
# Rotary RoPE Attention
# ============================================================================
class RoPE2DAttention(nn.Module):
"""
Attention with 2D Rotary Position Embedding applied to Q and K.
Uses standard RoPE formulation extended to 2D with adaptive aspect-ratio-aware positions.
Position mapping is normalized such that max(H) * max(W) = 1 and preserves aspect ratio.
"""
def __init__(self, hidden_size: int, num_heads: int, max_period: int = 10000, qk_norm: bool = False, bias: bool = True):
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
self.max_period = max_period
self.qk_norm = qk_norm
if qk_norm:
self.qk_scale = nn.Parameter(torch.full([num_heads, 1, 1], 10.0))
self.q_proj = nn.Linear(hidden_size, hidden_size, bias=bias)
self.k_proj = nn.Linear(hidden_size, hidden_size, bias=bias)
self.v_proj = nn.Linear(hidden_size, hidden_size, bias=bias)
self.out_proj = nn.Linear(hidden_size, hidden_size, bias=bias)
# Precompute frequency bands (quarter for h, quarter for w)
quarter_head_dim = self.head_dim // 4
freqs = 1.0 / (
max_period
** (
torch.arange(0, quarter_head_dim, dtype=torch.float32)
/ quarter_head_dim
)
)
self.register_buffer("freqs", freqs)
def _compute_adaptive_rope_freqs(self, h: int, w: int, device: torch.device):
"""Compute RoPE frequencies for 2D grid with adaptive aspect-ratio mapping."""
grid_h, grid_w = get_adaptive_2d_pos(h, w, device)
freqs_h = torch.outer(grid_h, self.freqs.to(device)) # [L, head_dim/4]
freqs_w = torch.outer(grid_w, self.freqs.to(device)) # [L, head_dim/4]
freqs = torch.cat([freqs_h, freqs_w], dim=-1) # [L, head_dim/2]
return torch.cos(freqs), torch.sin(freqs)
def _apply_rope(
self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor
) -> torch.Tensor:
"""Apply rotary position embedding."""
# x: [B, num_heads, L, head_dim]
# cos, sin: [L, head_dim/2]
x_half = x.shape[-1] // 2
x1, x2 = x[..., :x_half], x[..., x_half:]
cos = cos.unsqueeze(0).unsqueeze(0) # [1, 1, L, head_dim/2]
sin = sin.unsqueeze(0).unsqueeze(0)
return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1)
@compile_wrapper
def forward(self, x: torch.Tensor, h: int, w: int) -> torch.Tensor:
"""
Args:
x: [B, T, S, D] input (T=temporal, S=sequence length)
S may be h*w (spatial only) or h*w + extra tokens (motion, timestep, etc.)
h, w: spatial dimensions (h*w <= S)
"""
B, T, S, D = x.shape
L = h * w # spatial token count
total = T * S
x_flat = x.flatten(1, 2) # [B, T*S, D]
q = self.q_proj(x_flat)
k = self.k_proj(x_flat)
v = self.v_proj(x_flat)
q = q.view(B, total, self.num_heads, self.head_dim).transpose(1, 2)
k = k.view(B, total, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(B, total, self.num_heads, self.head_dim).transpose(1, 2)
# Get adaptive RoPE frequencies for spatial positions
cos_spatial, sin_spatial = self._compute_adaptive_rope_freqs(h, w, x.device)
# cos_spatial, sin_spatial: [L, head_dim/2]
# Build per-token RoPE: spatial tokens get position encoding,
# extra tokens (motion, timestep) get zero rotation (cos=1, sin=0)
half = cos_spatial.shape[-1]
if S > L:
extra = S - L
ones = torch.ones(extra, half, device=x.device, dtype=cos_spatial.dtype)
zeros = torch.zeros(extra, half, device=x.device, dtype=cos_spatial.dtype)
cos_frame = torch.cat([cos_spatial, ones], dim=0) # [S, half]
sin_frame = torch.cat([sin_spatial, zeros], dim=0) # [S, half]
else:
cos_frame = cos_spatial
sin_frame = sin_spatial
# Tile for all temporal frames
cos = cos_frame.unsqueeze(0).expand(T, -1, -1).reshape(total, -1)
sin = sin_frame.unsqueeze(0).expand(T, -1, -1).reshape(total, -1)
# Apply RoPE to Q and K
q = self._apply_rope(q, cos, sin)
k = self._apply_rope(k, cos, sin)
# QK-norm (after RoPE)
if self.qk_norm:
from .layers import _qk_norm
q, k = _qk_norm(q, k, self.qk_scale)
attn = F.scaled_dot_product_attention(q, k, v, scale=1.0)
else:
attn = F.scaled_dot_product_attention(q, k, v)
out = attn.transpose(1, 2).flatten(2, 3) # [B, T*S, D]
out = self.out_proj(out)
return out.view(B, T, S, D)