# 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 [--ckpt_path=...]` in one tmux; `scripts/_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](https://github.com/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_/`, 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/_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 --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 ```