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)
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support