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 S3s3://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 vseval_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. Statictmuxat~/.local/bin/tmux. - W&B: project
pixeldit-imagenet-reproonadobesensei.wandb.io(API key needed;wandb==0.19.11required β 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.shin another (WANDB_API_KEY,PIXELDIT_SKIP_S3_RESTORE=1set). 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 lossy4489tgd, FIDxo84hpbc. - 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&Bs075fe0g/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&B0e4hcesw/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.SRATrainerinsrc/diffusion.py(use_projflag),pix256_xl_sra.yaml. W&B77xtrzcq/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&B5ywf06q4/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 run49p8qfc5; 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
- 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.
- 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).
- 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.
- DiT-blk12 as teacher (B) < DINOv2 (A); dino+VAE (C) β baseline (slightly better late).
- Diffusion-model features need DC removal to be a usable cosine-alignment teacher (B, C).
- 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 sessionshybrid_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.shcan 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; setSTEPS=...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