MAVT / MODEL.md
Anbinh93's picture
Initial upload: code + configs + Stage 3 live progress (rgat-demo branch)
251713e verified
|
Raw
History Blame Contribute Delete
17 kB

MAVT (AnTokenizer) — Model & Pipeline

Unified vision tokenizer cho image / video / 3D, học chung 1 latent space. Train pipeline 3-stage curriculum (image → +video → +3D). Code chính: src/mavt/model/ + src/mavt/training/lightning_module.py.


1. Pipeline tổng quan (7 stages)

                 ┌─────────────────┐
   raw input ──▶ │ 1. Patchify     │  Conv3d, modality-specific
                 │   (multi-modal) │  → tokens (B,N,D), positions (N,4), plane_ids (N,)
                 └────────┬────────┘
                          │
                 ┌────────▼────────┐
                 │ 2. Hybrid       │  Transformer + RGAT, 12 blocks, dim=768
                 │   Backbone      │  → features (B,N,D)
                 └────────┬────────┘
                          │
                 ┌────────▼────────┐
                 │ 3. C-D Split    │  GLOBAL content (slot attn) + LOCAL detail (window pool)
                 │   (cd-split)    │  → compressed (B, N_c+N_d, D), latent_positions, token_types
                 └────────┬────────┘
                          │
                 ┌────────▼────────┐
                 │ 4. VAE bottleneck│  N×D → N×latent_dim (32), reparametrize z
                 │   + KL loss     │  → z, mu, logvar, loss_kl
                 └────────┬────────┘
                          │
              ┌───────────┴───────────┐
       ┌──────▼──────┐         ┌──────▼──────┐
       │ 5a. Recon   │         │ 5b. Underst.│
       │   Decoder   │         │   Decoder   │
       │ (z → pixel) │         │ (z → semantic, distill từ SigLIP2)
       └──────┬──────┘         └──────┬──────┘
              │                        │
              ▼                        ▼
        recon (B,3,H,W)         semantic (B, 768)
              │                        │
              └─────────┬──────────────┘
                        │
              ┌─────────▼─────────┐
              │ 6. Loss (MAVTLoss)│  l1 + lpips + KL + sem distill + temporal (video) + slot div
              │   per-modality EMA│  proportional weighting
              └─────────┬─────────┘
                        │
                        ▼
              7. Outputs (`MAVTOutput` dataclass)

2. Stage chi tiết

Stage 1 — Patchify (patchify.py)

  • Image (B, 3, H, W): Conv3d với causal-pad 1 frame ảo → (B, D, 1, Hp, Wp) → flatten (B, Hp·Wp, D). Position (0, i, j, 0), plane_id=−1.
  • Video (B, 3, T, H, W): Conv3d giảm temporal t_patch=2, spatial patch=16(B, D, Tp, Hp, Wp)(B, Tp·Hp·Wp, D). Position (t, i, j, 0), plane_id=−1.
  • Threed (B, 3, 3, S, S) (3 planes XY/XZ/YZ): Conv3d riêng cho mỗi plane → concat (B, 3·Hp·Wp, D). Position encode plane-specific axes; plane_id ∈ {0,1,2}.
  • Tất cả share Conv3d weight → unified patch embedding qua modality.
  • Kèm 4D pos_embed (Fourier features, dim D) cộng vào tokens.

Stage 2 — Hybrid Backbone (backbone.py)

  • 12 layers xen kẽ Transformer blocks + RGAT (Relational Graph Attention).
  • RGAT là attention dạng graph với:
    • r_s=2 neighborhoods spatial (window-based)
    • r_t=1 neighborhood temporal
  • use_gradient_checkpointing=true: tradeoff compute/memory (đang OFF cho run tuned để tăng tốc).
  • SigLIP2 weights được load vào last 4 transformer blocks (Stage 1: frozen, Stage 2: unfrozen).

Stage 3 — Content-Detail Split ⭐ (deep dive section 3)

Stage 4 — VAE bottleneck (latent_heads.py)

  • VAEHead: Linear(D=768 → 2·latent_dim=64) → split mu, logvar → reparametrize trick.
  • KL loss: KL(N(mu, σ²) || N(0, I)), scale bằng kl_weight=1e-4 (built-in).
  • Output z shape (B, N_c+N_d_local, latent_dim=32).
  • Note: post 42c622b update, w_kl=1.0 ở MAVTLoss = passthrough — đã pre-scale 1e-4 trong VAEHead → tránh double-scale.

Stage 5 — Decoders (decoder.py)

5a. Reconstruction AsymmetricDecoder:

  • UnifiedDetailExpander cross-attention từ target grid positions vào latent z (dim 32).
  • 4 self-attention blocks dim 768.
  • Pixel projection → (B, 3, H, W) cho image, frame-by-frame cho video.
  • Mới (cd-split): nhận latent_positions + latent_token_types để áp distance-bias attention (xem section 3).

5b. Understanding UnderstandingDecoder:

  • 2 cross-attn layers + linear proj → (B, semantic_dim=768).
  • Trained để khớp với SigLIP2 teacher's pooler_output qua cosine loss.

Stage 6 — Loss (losses.py)

L_total = w_l1 · L1(pred, target)
        + w_lpips · LPIPS(pred, target)         # AlexNet/VGG perceptual
        + w_kl · L_KL                            # đã pre-scaled ở VAEHead
        + w_sem · (1 - cos(MAVT.semantic, teacher.pooler))
        + w_temp · L1(Δ_t pred, Δ_t target)     # chỉ video, T>1
        + w_aux · slot_diversity_penalty
  • Mỗi modality scale bằng ModalityEMAWeighter.weight(modality):
    • weight(m) = ema_m / mean(ema_active) — modality có loss CAO được boost (cross-stage gradient flow vào branch chưa train)
    • Sau commit 42c622b update. Trước đó là 1/ema (logic ngược).

3. ⭐ C-D Split (sau commit 6368dfb cd-split)

Ý tưởng cốt lõi

Phân chia input tokens thành 2 kênh có đặc tính khác nhau, encode bằng cơ chế khác nhau:

Kênh Bản chất Cơ chế Position info
Content Semantic / low-freq / global Slot cross-attention (toàn ảnh) Không có (slot là global summary)
Detail Residual / high-freq / local Coordinate window pooling (window center)

Tại sao Detail cần local + position?

Trước cd-split (slot pooler global cho cả detail):

detail = SlotPooler(N_d=25 slots)(Residual)
       (slots tự học pool ở đâu, không có vị trí)
  • Decoder cross-attend vào detail slots không biết slot này ứng với patch nào → phải reconstruct texture từ "positionless global slots" → khó.
  • High-freq (texture, edges) cần spatial precision → mất khi pool global.

Sau cd-split (windowed pool + position):

# Group residual tokens theo coordinate window
group_key = (plane_id, t // t_win, x // s_win, y // s_win, z // s_win)
# Mean-pool tokens trong cùng window
detail_token[g] = mean(residual[token] for token in window g)
detail_position[g] = mean(positions[token] for token in window g) + 0.5 # window center
  • Mỗi detail token có toạ độ rõ ràng → decoder biết detail thuộc patch nào.
  • Decoder dùng distance bias (Manhattan) trong cross-attn để mỗi pixel ưu tiên detail token gần.
  • Compression vẫn tốt: window 2×2 → 4 token → 1 token (75% giảm), tổng compression vẫn ~50% (bằng N_c + N_d_local).

Architecture sau update

                    ┌─ slot attn ──▶  C  (B, N_c, D)        [global, positionless]
features ──┬─▶ ─────┤
(B,N,D)    │        └─ approx via inverse softmax weights:
           │           x_approx = softmax(C @ xᵀ / √D)ᵀ @ C
           │
           └─▶ R = x - x_approx (residual)
                    │
                    ▼
              ┌──────────────────────────────────────┐
              │   _local_detail_pool(R, pos, plane)  │
              │                                       │
              │  group_key = (plane, t/1, i/2, j/2, k/2)
              │  D_tokens[g] = mean(R[t] for t in g)  │
              │  D_pos[g]    = mean(pos[t]) + 0.5     │
              │  D_tokens   ← detail_proj(detail_norm(.))
              └──────────┬───────────────────────────┘
                         │
                         ▼
                    detail tokens (B, N_d_local, D)  + detail_positions (N_d_local, 4)

compressed = concat([C, detail_tokens])  (B, N_c + N_d_local, D)
latent_positions  = concat([zeros(N_c, 4), detail_positions])
latent_token_types = concat([zeros(N_c), ones(N_d_local)])

Decoder sử dụng metadata thế nào?

# UnifiedDetailExpander forward
kv = z + kv_pos_scale * kv_pos_enc(latent_positions)        # add 4D Fourier pos
       + token_type_scale * token_type_embed(latent_token_types)  # +0/+1 embedding

# Distance bias chỉ apply cho detail keys
dist_manhattan = |query_pos - kv_pos|
attn_bias[detail_keys] = -local_detail_bias * dist  # local_detail_bias=0.25

cross_attn(query, kv, kv, attn_mask=attn_bias)

→ Pixel position xa detail position thì attention bị penalize logarithmically (softmax-scale). → Content tokens KHÔNG bị penalize → decoder vẫn dùng được toàn bộ semantic info.

Worked example: image 256×256

Input: x  shape = (B, 3, 256, 256)
Patch_size=16  →  Hp=Wp=16  →  N=256 tokens
positions = [(0,i,j,0) for i,j in 16×16]  → (256, 4)
plane_ids = [-1] * 256

content_ratio=0.25, detail_ratio=0.25 → N_c=64, N_d_key=64 (key naming only)

Stage 3a: Content slot pool

slot_pooler = SlotPooler(num_slots=64, dim=768, num_heads=8, num_layers=2)
C = slot_pooler(features)  # 2 cross-attn layers
   shape (B, 64, 768)
  • 64 learned slots cross-attend toàn 256 tokens → mỗi slot là weighted summary toàn ảnh.

Stage 3b: Approximate + residual

weights = softmax(C @ features.T / sqrt(768), dim=-1)  # (B, 64, 256)
x_approx = weights.T @ C                                # (B, 256, 768)
R = features - x_approx                                 # (B, 256, 768)  — high-freq

Stage 3c: Local detail pool (window=2)

group_key[token_n] = (plane_id=-1, t=0, i//2, j//2, z=0)
                   = (-1, 0, i//2, j//2, 0)

i=0,j=0 → key (-1,0,0,0,0)   group 0
i=0,j=1 → key (-1,0,0,0,0)   group 0  (same window 2×2)
i=0,j=2 → key (-1,0,0,1,0)   group 1
i=0,j=3 → key (-1,0,0,1,0)   group 1
...
i=1,j=0 → key (-1,0,0,0,0)   group 0
i=1,j=1 → key (-1,0,0,0,0)   group 0
...

→ 4 token (vd i=0..1, j=0..1) gộp vào group 0. → Tổng số group = 8×8 = 64 detail tokens.

counts[0] = 4 (i=0,1; j=0,1)
D_token[0] = mean(R[0], R[1], R[16], R[17])     # 4 token trong window 2×2
D_token[0] = detail_proj(detail_norm(D_token[0]))

D_pos[0] = mean([(0,0,0,0), (0,0,1,0), (0,1,0,0), (0,1,1,0)]) + 0.5 = (0, 1, 1, 0)
                                                                  → window center floor

Output:

  • compressed = concat(C, D_tokens) shape (B, 128, 768)
  • latent_positions shape (128, 4):
    • First 64 rows: (0,0,0,0) (content, positionless)
    • Last 64 rows: window centers like (0, 1, 1, 0), (0, 1, 3, 0), ...
  • latent_token_types shape (128,): [0]*64 + [1]*64

Worked example: video 256×256, 16 frames

T=16, t_patch=2 → Tp=8
N = 8 × 16 × 16 = 2048 tokens

content_ratio=0.25 → N_c=512
detail_ratio=0.25  → N_d_key=512 (naming)

Detail windows (s_win=2, t_win=1):
  group_key = (plane=-1, t//1=t, i//2, j//2, 0)
  t in [0..7]:    8 unique
  i//2 in [0..7]: 8 unique
  j//2 in [0..7]: 8 unique
  → 8 × 8 × 8 = 512 detail windows
  • Compressed: 512 + 512 = 1024 tokens (vs 2048 raw → 2× compression)
  • Mỗi detail token gồm 1 temporal × 4 spatial residual (tổng 4 raw tokens).

Compare số token: image 256² (3 modality khác nhau)

Modality N raw N_c (content) N_d_local (detail, win=2) Total Compression
Image 256 64 64 128
Video 2048 512 512 1024
Threed 768 268 192 460 1.67×

Hyperparameters (configurable qua CLI hoặc yaml)

Param Default Tác động
local_detail_window_size 1 (sau user update) Kích thước window spatial. 1 = không pool (mỗi token 1 group), 2 = 2×2 windows
local_detail_temporal_window_size 1 Window temporal. 1 = mỗi frame riêng
content_ratio (modality-specific) 0.25 (img/vid), 0.35 (3D) N_c = N × ratio
detail_ratio (key naming only) 0.25 Không ảnh hưởng số detail token thực
local_detail_bias (decoder) 0.25 Hệ số distance bias trong cross-attn. Lớn = ép detail mạnh hơn
kv_pos_scale (decoder, learnable) init 0.1 Trọng số position encoding cộng vào KV
token_type_scale (decoder, learnable) init 0.1 Trọng số token type embed cộng vào KV

Monitoring metrics

Metric Ý nghĩa Target
slot_diversity mean pairwise cos sim giữa các content slots ≤ 0.5 (slots khác nhau)
residual_ratio ‖R‖ / ‖x‖ 0.3–0.5 (content giữ phần lớn signal)
detail_token_count số detail token thực = số window distinct
detail_avg_window_tokens trung bình tokens per window ≈ s_win² · t_win nếu density đều

Tại sao đổi từ slot pooler global → window pool cho detail?

Khía cạnh Slot pooler (cũ) Window pool (mới)
Position info ❌ (global) ✅ (window center)
Compression tốt (25 slots cho 256 token) tốt (64 windows cho 256 token với win=2)
High-freq detail hạn chế (slot abstract) tốt (mean trong window nhỏ giữ texture)
Inductive bias không có spatial prior có (locality assumption)
Tham số trainable slot params + 2 layer cross-attn non-trainable scatter_add + 1 LayerNorm + 1 Linear
Compute cao (slot attn O(N·N_d)) thấp (scatter O(N))

→ Chuyển sang windowed pool cho detail là lossless về expressive power với inductive bias hợp lý cho high-freq, lại rẻ hơn về tham số/compute.


4. Curriculum 3 stage

Stage Modalities SigLIP2 unfreeze LR Purpose
1 image only hoàn toàn frozen 1e-4 Học image tokenizer + distill semantic
2 image + video last 4 blocks unfrozen 5e-5 Thêm video poolers, fine-tune backbone cho temporal
3 + threed toàn bộ unfrozen 2e-5 Thêm 3D, polish toàn bộ

Cross-stage transfer: --model.init_from_ckpt <prev_stage_ckpt> (strict=False, weights only). Lightning module setup('fit'):

  1. _prepare_cd_split_poolers() — eagerly tạo content pooler cho mỗi modality active (bắt buộc trước configure_optimizers)
  2. _sync_ema_modalities() — đồng bộ active_modalities từ DataModule vào EMA weighter
  3. load_siglip2_weights() — load HF weights nếu init_siglip2=true
  4. _load_semantic_teacher() — load frozen SigLIP2 vision tower nếu use_semantic_distill=true
  5. _load_weights_from_ckpt(init_from_ckpt) — load prev-stage weights cuối cùng để override init

5. Latent space chính thức

Sau VAE bottleneck:

  • Image: 128 token × 32 dim = 4096 floats (vs raw 196,608 → 48× compression)
  • Video: 1024 token × 32 dim = 32,768 floats (vs raw 3,145,728 → 96× compression)
  • Threed: 460 token × 32 dim = 14,720 floats (vs raw 589,824 → 40× compression)

So với baselines (theo results.md):

  • AToken-So/C Stage 1: (1, 16, 16) = 256 token × 32 ch = 8192 floats / image (chúng tôi 4096 — gấp đôi compression nhờ slot ratio 0.25)
  • Cosmos-CI16×16: 256 × 16 = 4096 (same as ours per token volume, nhưng arch khác)

6. Reference files

File Content
src/mavt/model/patchify.py Stage 1
src/mavt/model/backbone.py Stage 2 (Transformer + RGAT)
src/mavt/model/content_detail_split.py Stage 3 ⭐
src/mavt/model/latent_heads.py Stage 4 (VAEHead)
src/mavt/model/decoder.py Stage 5 (AsymmetricDecoder, UnderstandingDecoder, UnifiedDetailExpander)
src/mavt/losses/losses.py Stage 6 (MAVTLoss, ModalityEMAWeighter, temporal_consistency_loss)
src/mavt/training/lightning_module.py Curriculum, optimizer, logging
src/mavt/model/mavt.py End-to-end MAVT module
configs/model/mavt_base.yaml Hyperparams arch
configs/train/universal_data/stage{1,2,3}_universal.yaml Stage curriculum