Files
enso/README.md
T
2024-07-18 01:20:54 -05:00

3.3 KiB
Raw Blame History

Scaling Diffusion Transformers with Mixture of Experts
Official PyTorch Implementation

arXiv

This repo contains PyTorch model definitions, pre-trained weights and training/sampling code for our paper scaling Diffusion Transformers to 16 billion parameters (DiT-MoE). DiT-MoE as a sparse version of the diffusion Transformer, is scalable and competitive with dense networks while exhibiting highly optimized inference.

DiT-MoE framework

  • 🪐 A PyTorch implementation of DiT-MoE
  • Pre-trained checkpoints in paper
  • 💥 A sampling script for running pre-trained DiT-MoE
  • 🛸 A DiT-MoE training script using PyTorch DDP and FSDP

To-do list

  • training / inference scripts
  • huggingface ckpts
  • experts routing analysis
  • synthesized data

1. Training

You can refer to the link to build the running environment.

To launch DiT-MoE-S/2 (256x256) in the latent space training with N GPUs on one node with pytorch DDP:

torchrun --nnodes=1 --nproc_per_node=N train.py \
--model DiT-S/2 \
--num_experts 8 \
--num_experts_per_tok 2 \
--data-path /path/to/imagenet/train \
--image-size 256 \
--global-batch-size 256 \
--vae-path /path/to/vae

For multiple node training, we solve the bug at original DiT repository, and you can run with 8 nodes as:

torchrun --nnodes=8 \
    --node_rank=0 \
    --nproc_per_node=8 \
    --master_addr="10.0.0.0" \
    --master_port=1234 \
    train.py \
    --model DiT-B/2 \
    --num_experts 8 \
    --num_experts_per_tok 2 \
    --global-batch-size 1024 \
    --data-path /path/to/imagenet/train \
    --vae-path /path/to/vae

2. Inference

We include a sample.py script which samples images from a DiT-MoE model.

python sample.py \
--model DiT-S/2 \
--ckpt /path/to/model \
--image-size 256 \
--cfg-scale 1.5

3. Download Models and Data

We are processing it as soon as possible, the model weights and data will be released within two weeks :)

DiT-MoE Model Image Resolution Url
DiT-MoE-S/2-8E2A 256x256 -
DiT-MoE-S/2-16E2A 256x256 -
DiT-MoE-B/2-8E2A 256x256 -
DiT-MoE-XL/2-8E2A 256x256 -
DiT-MoE-XL/2-8E2A 512x512 -
DiT-MoE-G/2-16E2A 512x512 -

4. Expert Specialization Analysis Tools

We provide all the analysis scripts used in the paper.
You can use expert_data.py to sample data points towards experts ids across different class-conditional. Then, the headmap.py is used to viasualize frequency for different scenarios.

5. BibTeX

@article{FeiDiTMoE2024,
  title={Scaling Diffusion Transformers to 16 Billion Parameters},
  author={Zhengcong Fei, Mingyuan Fan, Changqian Yu, Debang Li, Jusnshi Huang},
  year={2024},
  journal={arXiv preprint},
}

6. Acknowledgments

The codebase is based on the awesome DiT and DeepSeek-MoE repos.