mirror of
https://github.com/rohitg00/ai-engineering-from-scratch.git
synced 2026-10-02 01:54:39 +08:00
fix(phase-4/28): correct divided-attention complexity math (H*W)*T^2 + T*(H*W)^2 in docs, quiz, and main.py; drop unused math import
This commit is contained in:
@@ -1,4 +1,3 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
@@ -69,11 +68,17 @@ def count_tokens(T, H, W, p_t=2, p_h=8, p_w=8):
|
||||
def main():
|
||||
print("[token count for 5s 360p video (150 frames, 480x360)]")
|
||||
tokens = count_tokens(150, 480, 360, p_t=2, p_h=8, p_w=8)
|
||||
T_tok = 150 // 2
|
||||
S_tok = (480 // 8) * (360 // 8)
|
||||
print(f" tokens per clip: {tokens:,}")
|
||||
print(f" attention pairs (joint): {tokens ** 2:,}")
|
||||
print(f" divided time: {(150 // 2) ** 2:,}")
|
||||
print(f" divided space: {((480 // 8) * (360 // 8)) ** 2:,}")
|
||||
print(f" divided total: {(150 // 2) ** 2 + ((480 // 8) * (360 // 8)) ** 2:,}")
|
||||
# Divided temporal: T^2 attention at every spatial position.
|
||||
# Divided spatial: (H*W)^2 attention at every timestep.
|
||||
divided_time = S_tok * T_tok ** 2
|
||||
divided_space = T_tok * S_tok ** 2
|
||||
print(f" divided time total: {divided_time:,}")
|
||||
print(f" divided space total: {divided_space:,}")
|
||||
print(f" divided total: {divided_time + divided_space:,}")
|
||||
|
||||
torch.manual_seed(0)
|
||||
vid = torch.randn(1, 4, 8, 16, 16)
|
||||
|
||||
@@ -61,7 +61,7 @@ Resulting tokens: (T / P_t) * (H / P_h) * (W / P_w) tokens
|
||||
Positional encoding is 3D: a rotary or learned embedding per (t, h, w) coordinate. Attention can be:
|
||||
|
||||
- **Full joint** — all tokens attend to all tokens. O(N^2) with N tokens. Prohibitive for long videos.
|
||||
- **Divided** — alternate temporal attention (same spatial position, across time) and spatial attention (same timestep, across space). Used by TimeSformer and most video DiTs.
|
||||
- **Divided** — alternate temporal attention (same spatial position, across time: `(H*W) * T^2`) and spatial attention (same timestep, across space: `T * (H*W)^2`). Used by TimeSformer and most video DiTs.
|
||||
- **Window** — local windows in (t, h, w). Used by Video Swin.
|
||||
|
||||
Every 2026 video diffusion model uses one of these three patterns plus AdaLN conditioning (Lesson 23) and rectified flow.
|
||||
|
||||
@@ -10,9 +10,9 @@
|
||||
{
|
||||
"stage": "pre",
|
||||
"question": "Divided attention in a video transformer means what?",
|
||||
"options": ["Half the tokens are masked", "Each block does a temporal attention (same spatial position, across frames) followed by a spatial attention (same frame, across positions); this reduces cost from O((T*H*W)^2) to O(T^2) + O((H*W)^2)", "Only half the layers run attention", "Attention is split across GPUs"],
|
||||
"options": ["Half the tokens are masked", "Each block does a temporal attention (same spatial position, across frames) followed by a spatial attention (same frame, across positions); this factorises cost from O((T*H*W)^2) into O((H*W)*T^2) + O(T*(H*W)^2) — dramatically cheaper than the joint product", "Only half the layers run attention", "Attention is split across GPUs"],
|
||||
"correct": 1,
|
||||
"explanation": "Full joint attention over spacetime tokens is prohibitive: for 150 frames at 60x45 token grid you'd be doing a 405,000^2 attention. Divided attention alternates temporal (300^2) and spatial (2700^2) attention, factoring the complexity. TimeSformer introduced this pattern; almost every 2026 video DiT (Sora, Wan, HunyuanVideo) uses a divided or window variant."
|
||||
"explanation": "Full joint attention over spacetime tokens is prohibitive: for T=150 temporal tokens and a 60x45 spatial grid (2700 spatial tokens), the joint (T*H*W)^2 ≈ 1.6e11 pairs. Divided attention runs temporal attention at each spatial position (H*W * T^2 ≈ 6.1e7) and spatial attention at each timestep (T * (H*W)^2 ≈ 1.1e9) — multiple orders of magnitude less. TimeSformer introduced this pattern; almost every 2026 video DiT (Sora, Wan, HunyuanVideo) uses a divided or window variant."
|
||||
},
|
||||
{
|
||||
"stage": "post",
|
||||
|
||||
Reference in New Issue
Block a user