DSR β RAE-DiT-XL ImageNet-256 reference checkpoints
Stage-2 diffusion-transformer checkpoints accompanying DSR: Dual-Stage Registers (Taming Outlier Tokens in Diffusion Transformers), released so the outlier-token visualizations can be reproduced without retraining.
Both are EMA weights of a DiT-XL (depth 28, hidden 1152, 16 heads, patch 1,
in_channels 768) operating in the SigLIP2-B RAE latent space (16Γ16Γ768,
ImageNet-256 class-conditional), matching the DSR repo's
configs/stage2/sampling/ImageNet256/DiT-XL_SigLIP2-B_data_lognormal.yaml.
| File | Model | Training config (DSR repo) | Notes |
|---|---|---|---|
dit_xl_siglip2b_lognormal_ema.pt |
baseline, no registers | DiT-XL_SigLIP2-B_data_lognormal |
full training run |
dit_xl_siglip2b_lognormal_in_context_36_ep80_ema.pt |
+ 36 in-context diffusion registers at block 8 | *_data_lognormal-in_context_36 |
epoch-80 snapshot (early): samples look under-trained; intended for comparing token-norm structure, where the register effect is already fully developed |
Usage
The DSR notebook reproduce_outlier_tokens.ipynb
downloads these automatically when no local checkpoint is found. Manual loading:
import torch
from huggingface_hub import hf_hub_download
path = hf_hub_download("smallAI/DSR-DiT-XL-ImageNet256",
"dit_xl_siglip2b_lognormal_ema.pt")
state_dict = torch.load(path, map_location="cpu", weights_only=True)
# build LightningDiT from the DSR repo's sampling config, then:
# model.load_state_dict(state_dict, strict=True)
Both files are plain state_dicts (no optimizer state) and load with
torch.load(..., weights_only=True).
What they demonstrate
Sampling the baseline while recording per-token L2 norms at every block shows outlier tokens β max/median patch-token norm rising to ~20Γ in blocks ~14β24, at fixed spatial positions shared across samples. In the register model the patch tokens stay near-flat and the 36 register tokens carry the outlier mass instead. See the DSR repository for the full method and the notebook that produces these figures.
License / intended use
Released for research reproducibility. The models were trained on ImageNet-1k; use of these weights must comply with the ImageNet terms of access.
Stage-1 artifacts (for pixel-space decoding)
Mirrors of the stage-1 release linked from the DSR README (Google Drive), needed only for the notebook's whole-pipeline figure (encode β noise β DiT β decode):
| File | What |
|---|---|
rae_decoder_vitxl_siglip2b_256.pt |
RAE ViT-XL decoder for the SigLIP2-B latent space (ImageNet-256) |
rae_stats_siglip2b_imagenet1k.pt |
Per-channel latent normalization statistics (mean / var) |