Video Classification
Transformers
Safetensors
ttvidt
feature-extraction
video
video-representation-learning
self-supervised-learning
motion
temporal-modeling
dinov3
vision-transformer
custom_code
Eval Results (legacy)
Instructions to use KBlueLeaf/TTVidT with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use KBlueLeaf/TTVidT with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("video-classification", model="KBlueLeaf/TTVidT", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("KBlueLeaf/TTVidT", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download pos_embed.py from KBlueLeaf/TTVidT: direct link, hf CLI and curl.
- Browser
- Download file 12.6 kB
-
https://huggingface.co/KBlueLeaf/TTVidT/resolve/main/pos_embed.py
- Command line
-
hf download hf://KBlueLeaf/TTVidT/pos_embed.py
-
curl -L -o pos_embed.py https://huggingface.co/KBlueLeaf/TTVidT/resolve/main/pos_embed.py
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) | |
| 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) | |