repa / EXPERIMENTS_REPA.md
linxin02's picture
Update REPA metrics and experiment documentation
9a11bd8 verified
|
Raw
History Blame Contribute Delete
11.1 kB

PixelDiT Representation-Alignment (REPA) β€” experiments & handoff

Study of what to align PixelDiT's intermediate features to (the "teacher" in REPA-style training) on ImageNet-256, plus a per-layer feature analysis. Variants A–E keep PixelDiT's flow-matching loss and alignment point (patch block 8) fixed and only change the teacher; variant F is a follow-up hybrid that adds SRA guidance at block 4 while retaining block-8 REPA. This doc is the authoritative summary + handoff (results, infra, current state, how to continue).


1. Setup / infra

  • Model: PixelDiT-XL, 797M β€” pixel-space DiT, 26 patch blocks + 4 pixel blocks, patch 16, 256Γ—256.
  • Data: ImageNet-1K, REPA-E 256px H5 at imagenet256_data/images.h5 (253 GB, restored from S3 s3://dit-scale-up/fe-experiments/xinlin/pixeldit_imagenet256_xl_lx_vb/data/).
  • Recipe (identical across variants unless noted): 8Γ—80GB, global batch 256 (8Γ—32), AdamW lr 1e-4, EMA 0.9999 (SRA/Staged use 0.999, see Β§3D), bf16-mixed, 70k steps, checkpoint every 10k (save_top_k=-1, all kept, ~12.8 GB each).
  • Eval: main.py predict β†’ 50k samples (50/class, FlowDPMSolver 100-step, CFG 2.75, interval [0.1,0.9]) β†’ OpenAI guided-diffusion ADM evaluator vs eval_refs/VIRTUAL_imagenet256_labeled.npz β†’ FID / sFID / IS. FID worker computes one point per 10k checkpoint.
  • Envs: /mnt/localssd/pixeldit-venv (torch 2.5.0+cu124, training/predict); ADM evaluator in /mnt/localssd/adm-eval-venv (tf 2.15) + /mnt/localssd/guided-diffusion. Static tmux at ~/.local/bin/tmux.
  • W&B: project pixeldit-imagenet-repro on adobesensei.wandb.io (API key needed; wandb==0.19.11 required β€” 40-char check rejects the self-hosted key on β‰₯0.20).
  • Run a variant (8-GPU): scripts/train_pixeldit_imagenet256_xl.sh <config.yaml> [--ckpt_path=...] in one tmux; scripts/<variant>_fid_eval.sh in another (WANDB_API_KEY, PIXELDIT_SKIP_S3_RESTORE=1 set). FID/predict shares the 8 GPUs with training (~0.06 it/s, ~1.5–2h per 50k point).

2. Per-layer feature analysis (which layer is semantic)

On the official converged PixelDiT-XL (nvidia/PixelDiT-ImageNet epoch320), hooking each patch block and PCAβ†’RGB + linear-CKA-vs-DINOv2: semantic content peaks sharply at patch block 8 (CKA 0.85), exactly the REPA anchor; early blocks = detail, late blocks specialize toward the denoising output. DiT-XL/2-256's own features are semantic mid-late (blk 12–21). Scripts: scripts/viz_layer_features.py, viz_layer_similarity.py, viz_dit_layers.py (figures in viz_layer_features/, viz_dit_layers/ β€” PNGs are gitignored, regenerate from scripts).

3. Experiments (teacher = the only thing that changes)

  • A. baseline (DINOv2 REPA) β€” project block-8 tokens (MLPβ†’768), cosine-align to frozen DINOv2 patch tokens of the clean image, weight 0.5. pix256_xl_orig_10k.yaml. W&B loss y4489tgd, FID xo84hpbc.
  • B. DiT-REPA (blk12) β€” teacher = frozen DiT-XL/2-256: VAE-encode β†’ noise at ddpm t=100 β†’ DiT forward (class-cond) β†’ block-12 features. Gotcha: raw DiT features are DC-dominated (inter-token cosine 0.965) so cosine aligns trivially β†’ fixed by per-image mean-centering. src/dit_repa.py, pix256_xl_ditrepa12.yaml. W&B s075fe0g / o113h6qc.
  • C. dino+VAE β€” keep DINOv2 REPA and add a 2nd projector on block 8 aligned to the frozen SD-VAE encoder mid-block feature (DC-centered), Ξ»_gen 0.05. src/dino_vae_repa.py, pix256_xl_dinovae.yaml. W&B 0e4hcesw / pixeldit_dinovae_fid.
  • D. SRA (self-distillation, EMA teacher) β€” no external teacher: student block-8 (through a projector) aligns to the EMA net's block-16 at a lower-noise timestep, smooth-L1 Ξ²0.05, weight 0.04. Port of vvvvvjdy/SRA. Two fixes were required: (1) use_proj=True β€” reference projects the STUDENT and aligns to the RAW teacher (asymmetric); (2) EMA decay 0.9999 β†’ 0.999 β€” SRA reads the EMA forward every step and drives the online weight-norm up; the wide 0.9999 window lags that norm so the EMA (not the online model) samples as noise by ~30k. A tighter 0.999 window tracks it β†’ clean EMA. SRATrainer in src/diffusion.py (use_proj flag), pix256_xl_sra.yaml. W&B 77xtrzcq / pixeldit_sra_fid.
  • E. StagedTeacher (self-distillation, frozen-snapshot teacher) β€” like SRA but the teacher is a frozen hard copy of the online model, refreshed every 10k steps (a periodically-updated target network) instead of an EMA. Resumed from SRA's 10k checkpoint. src/staged_teacher.py, pix256_xl_stageteach.yaml (refresh_interval: 10000). W&B 5ywf06q4 / pixeldit_stageteach_fid.
  • F. SRA4 + DINO-REPA8 (hybrid) β€” one online forward supplies two complementary targets: online block 4 (student projector) aligns by smooth-L1 to the EMA net's raw block 16 at a lower-noise timestep, while online block 8 independently aligns by cosine to frozen DINOv2 clean-image patch tokens. Weights remain SRA 0.04 and REPA 0.5; EMA decay is 0.999. src/hybrid_sra_repa.py, pix256_xl_sra_repa4.yaml. W&B training run 49p8qfc5; FID curve pending.

4. Results β€” 50k-sample ADM eval

FID (↓ better)

step baseline DiT-REPA dino+VAE SRA (EMA) StagedTeacher
10k 389.3 368.7 383.1 187.5 187.5ΒΉ
20k 363.6 350.6 387.5 102.8 116.0
30k 106.1 108.1 127.9 82.2 86.9
40k 66.8 81.7 52.6 67.1 68.7
50k 31.1 57.2 39.2 55.8 58.4
60k 21.3 39.6 25.9 β€” 49.5
70k β€” β€” 17.0 β€” 43.1

sFID (↓ better)

step baseline DiT-REPA dino+VAE SRA StagedTeacher
10k 414.7 417.3 398.4 50.4 50.4ΒΉ
20k 304.7 197.4 403.2 13.9 17.5
30k 78.2 45.6 73.2 10.0 13.0
40k 46.0 42.7 30.5 8.5 8.9
50k 20.3 32.2 27.0 7.9 7.8
60k 14.7 20.3 18.0 β€” 7.3
70k β€” β€” 14.1 β€” 6.8

IS (↑ better)

step baseline DiT-REPA dino+VAE SRA StagedTeacher
10k 1.07 1.04 1.06 5.40 5.40ΒΉ
20k 1.40 1.64 1.00 9.59 8.43
30k 11.71 10.48 9.94 11.15 10.39
40k 17.51 10.82 19.19 12.59 12.40
50k 26.65 13.84 24.78 14.19 13.73
60k 32.94 18.56 30.52 β€” 15.36
70k β€” β€” 37.30 β€” 17.12

ΒΉ StagedTeacher resumed from SRA's 10k checkpoint, so its 10k = SRA's 10k. Chart: slides_assets/fid_comparison_5way.png (regenerate; PNGs gitignored).

5. Findings

  1. Self-distillation (SRA, StagedTeacher) converges dramatically faster early and dominates sFID at every step. At 10k SRA is FID 187 / IS 5.4 while the external-teacher variants are still ~380 / IS ~1; SRA's 20k FID (103) β‰ˆ the DINOv2 baseline's 30k. sFID stays 2–5Γ— better throughout.
  2. But self-distillation plateaus on FID/IS. By 50–60k the external-teacher variants overtake on FID (60k: baseline 21.3, dino+VAE 25.9 vs StagedTeacher 49.5) and on IS. Net: self-distillation = fast + best structure (sFID), external semantic teacher = better final FID/diversity (IS). dino+VAE reaches the best FID by 70k (17.0).
  3. EMA teacher (SRA) β‰ˆ frozen-snapshot teacher (StagedTeacher), EMA slightly better early. Gap 20k 116 vs 103, narrowing to ~identical by 40k. The gain comes from the self-distillation mechanism (block-8β†’block-16 alignment), not from how the teacher is updated.
  4. DiT-blk12 as teacher (B) < DINOv2 (A); dino+VAE (C) β‰ˆ baseline (slightly better late).
  5. Diffusion-model features need DC removal to be a usable cosine-alignment teacher (B, C).
  6. Debugging lesson (SRA): the EMA "noise" was not EMA corruption β€” the EMA tracked online (cosine 0.985, drift 0) but lagged in norm under the wide 0.9999 window; several wrong fixes (stream, compile, clobber) were tried before the real cause (norm-lag) β†’ decay 0.999.

6. Current state (as of 2026-07-28)

  • Running: hybrid SRA4 + DINO-REPA8 (49p8qfc5), newly started from scratch. tmux sessions hybrid_train, hybrid_fid; the FID worker targets every 10k checkpoint from 10k through 70k.
  • Completed: StagedTeacher reached 70k; its final point is FID 43.11 / sFID 6.78 / IS 17.12.
  • Checkpoints on local SSD (train_logs/exp_pixeldit_imagenet256_xl_<variant>/, every 10k, ~12.8GB): baseline 10–60k, DiT-REPA 10–70k, dino+VAE 10–70k, SRA 10–60k, StagedTeacher 20–70k. NOT backed up to S3 β€” a box wipe loses them (scripts/sync_pixeldit_to_s3_loop.sh can back up).
  • baseline & SRA runs were stopped early (~62k) to free GPUs for later variants, hence FID only to 60k/50k.

7. Handoff notes (for continuing)

  • To finish a curve: relaunch its scripts/<variant>_fid_eval.sh (idempotent β€” skips computed points; set STEPS=... to target specific steps). Checkpoints are all saved, so any FID point can be (re)computed offline without retraining.
  • To resume/continue training: train_pixeldit_imagenet256_xl.sh <config> --ckpt_path=<...ckpt>.
  • Open directions: (a) push SRA/Staged past 70k or with a stronger Ξ» / IS-oriented tweak to fix the FID/IS plateau while keeping the sFID lead; (b) hybrid: DINOv2 REPA + SRA self-distillation (semantic teacher for late FID/IS + self-distillation for early speed & sFID); (c) tune SRA student/teacher blocks & sra_t_max. (d) back up checkpoints to S3 first.
  • Gotchas: EMA decay must be ≀0.999 for any variant that reads the EMA during training (norm-lag); keep DC-centering for diffusion/VAE-feature teachers; nn.Module.compile() is lazy (first-forward); the CLI list-index override --trainer.callbacks.N... is rejected by jsonargparse (edit the YAML).

8. File map

pixeldit_configs/pix256_xl_orig_10k.yaml    # A baseline (DINOv2)
pixeldit_configs/pix256_xl_ditrepa12.yaml   # B DiT-REPA blk12 (DC-centered)
pixeldit_configs/pix256_xl_dinovae.yaml     # C DINOv2 + SD-VAE detail head
pixeldit_configs/pix256_xl_sra.yaml         # D SRA (use_proj=true, EMA decay 0.999)
pixeldit_configs/pix256_xl_stageteach.yaml  # E StagedTeacher (refresh_interval 10000)
pixeldit_configs/pix256_xl_sra_repa4.yaml   # F hybrid: SRA 4->16 + DINOv2 REPA on block 8
PixelDiT/c2i/src/diffusion.py               # REPATrainer, DINOv2, SRATrainer (use_proj flag)
PixelDiT/c2i/src/dit_repa.py                # B
PixelDiT/c2i/src/dino_vae_repa.py           # C
PixelDiT/c2i/src/staged_teacher.py          # E
PixelDiT/c2i/src/hybrid_sra_repa.py          # F
scripts/*_fid_eval.sh                       # per-variant FID-curve workers
scripts/log_fid_wandb.py                    # append an FID point to a W&B run
scripts/viz_*_layers*.py, build_slides.py   # analysis + slide deck