sony/mf-rae

Python

40

3 commits

updated Nov 19, 2025

See the code

README

MeanFlow Transformers with Representation Autoencoders (RAE)
Official PyTorch Implementation

Paper

This repository is based on:

Environment

Dependency Setup

  1. Create environment and install via uv:
    conda create -n rae python=3.10 -y
    conda activate rae
    pip install uv
    
    # Install PyTorch 2.2.0 with CUDA 12.1
    uv pip install torch==2.2.0 torchvision==0.17.0 torchaudio --index-url https://download.pytorch.org/whl/cu121
    
    # Install other dependencies
    uv pip install timm==0.9.16 accelerate==0.23.0 torchdiffeq==0.2.5 wandb
    uv pip install "numpy<2" transformers einops omegaconf
    

Data Preparation

  1. Download ImageNet-1k raw data without preprocessing.
  2. Point Stage 1 and Stage 2 scripts to the training split via --data-path.

Pre-Training: Flow Matching

The RAE authors release flow-matching pre-traine models: RAE decoders, DiTDH diffusion transformers and stats for latent normalization. To download all models at once:


cd RAE
pip install huggingface_hub
hf download nyu-visionx/RAE-collections \
  --local-dir models 

To download specific models, run:

hf download nyu-visionx/RAE-collections \
  <remote_model_path> \
  --local-dir models 

Consistent Mid-Training

bash CMT_256.sh
bash CMT_512.sh

MFT and MFD Post-Training

Make sure to input the CMT checkpoint path obtained from the previous stage.

For instance, on ImageNet 512, they are

bash MFT_512.sh
bash MFD_512.sh

Distributed sampling for evaluation

Make sure to input the MeanFlow-RAE checkpoint path after training to the config file.

We provide our trained MF-RAE on Google Drive: https://drive.google.com/drive/folders/1EYVyIDKRZeHn6NO7uF5aJ1ycR3lvfnJu?usp=drive_link

bash Sample_256.sh
bash Sample_512.sh

Evaluation

ADM Suite FID setup

Use the ADM evaluation suite to score generated samples:

  1. Clone the repo:

    git clone https://github.com/openai/guided-diffusion.git
    cd guided-diffusion/evaluation
    
  2. Create an environment and install dependencies:

    conda create -n adm-fid python=3.10
    conda activate adm-fid
    pip install 'tensorflow[and-cuda]'==2.19 scipy requests tqdm
    
  3. Download ImageNet statistics (256×256 shown here):

    wget https://openaipublic.blob.core.windows.net/diffusion/jul-2021/ref_batches/imagenet/256/VIRTUAL_imagenet256_labeled.npz
    
  4. Evaluate:

    python evaluator.py VIRTUAL_imagenet256_labeled.npz /path/to/samples.npz
    

Acknowledgement

This code is built upon the following repositories:

  • SiT - for diffusion implementation and training codebase.
  • DDT - for some of the DiTDH implementation.
  • LightningDiT - for the PyTorch Lightning based DiT implementation.
  • MAE - for the ViT decoder architecture.
  • RAE - for the RAE model and checkpoints.

sony/mf-rae

Python

40

3 commits

updated Nov 19, 2025

See the code

README

MeanFlow Transformers with Representation Autoencoders (RAE)
Official PyTorch Implementation

Paper

This repository is based on:

Environment

Dependency Setup

  1. Create environment and install via uv:
    conda create -n rae python=3.10 -y
    conda activate rae
    pip install uv
    
    # Install PyTorch 2.2.0 with CUDA 12.1
    uv pip install torch==2.2.0 torchvision==0.17.0 torchaudio --index-url https://download.pytorch.org/whl/cu121
    
    # Install other dependencies
    uv pip install timm==0.9.16 accelerate==0.23.0 torchdiffeq==0.2.5 wandb
    uv pip install "numpy<2" transformers einops omegaconf
    

Data Preparation

  1. Download ImageNet-1k raw data without preprocessing.
  2. Point Stage 1 and Stage 2 scripts to the training split via --data-path.

Pre-Training: Flow Matching

The RAE authors release flow-matching pre-traine models: RAE decoders, DiTDH diffusion transformers and stats for latent normalization. To download all models at once:


cd RAE
pip install huggingface_hub
hf download nyu-visionx/RAE-collections \
  --local-dir models 

To download specific models, run:

hf download nyu-visionx/RAE-collections \
  <remote_model_path> \
  --local-dir models 

Consistent Mid-Training

bash CMT_256.sh
bash CMT_512.sh

MFT and MFD Post-Training

Make sure to input the CMT checkpoint path obtained from the previous stage.

For instance, on ImageNet 512, they are

bash MFT_512.sh
bash MFD_512.sh

Distributed sampling for evaluation

Make sure to input the MeanFlow-RAE checkpoint path after training to the config file.

We provide our trained MF-RAE on Google Drive: https://drive.google.com/drive/folders/1EYVyIDKRZeHn6NO7uF5aJ1ycR3lvfnJu?usp=drive_link

bash Sample_256.sh
bash Sample_512.sh

Evaluation

ADM Suite FID setup

Use the ADM evaluation suite to score generated samples:

  1. Clone the repo:

    git clone https://github.com/openai/guided-diffusion.git
    cd guided-diffusion/evaluation
    
  2. Create an environment and install dependencies:

    conda create -n adm-fid python=3.10
    conda activate adm-fid
    pip install 'tensorflow[and-cuda]'==2.19 scipy requests tqdm
    
  3. Download ImageNet statistics (256×256 shown here):

    wget https://openaipublic.blob.core.windows.net/diffusion/jul-2021/ref_batches/imagenet/256/VIRTUAL_imagenet256_labeled.npz
    
  4. Evaluate:

    python evaluator.py VIRTUAL_imagenet256_labeled.npz /path/to/samples.npz
    

Acknowledgement

This code is built upon the following repositories:

  • SiT - for diffusion implementation and training codebase.
  • DDT - for some of the DiTDH implementation.
  • LightningDiT - for the PyTorch Lightning based DiT implementation.
  • MAE - for the ViT decoder architecture.
  • RAE - for the RAE model and checkpoints.

Significant stargazers

Yuiga Wada

71 followers · starred Apr 2026