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 tt3d.py from KBlueLeaf/TTVidT: direct link, hf CLI and curl.
- Browser
- Download file 19.4 kB
-
https://huggingface.co/KBlueLeaf/TTVidT/resolve/main/tt3d.py
- Command line
-
hf download hf://KBlueLeaf/TTVidT/tt3d.py
-
curl -L -o tt3d.py https://huggingface.co/KBlueLeaf/TTVidT/resolve/main/tt3d.py
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 | |
| 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!") | |