The implementation of the paper "Not All Prompts Are Made Equal: Prompt-based Pruning of Text-to-Image Diffusion Models"
APTP: We prune a text-to-image diffusion model like Stable Diffusion (left) into a mixture of efficient experts (right) in a prompt-based manner. Our prompt router routes distinct types of prompts to different experts, allowing experts' architectures to be separately specialized by removing layers or channels.
APTP pruning scheme. We train the prompt router and the set of architecture codes to prune a T2I diffusion model into a mixture of experts. The prompt router consists of three modules. We use a Sentence Transformer as the prompt encoder to encode the input prompt into a representation z. Then, the architecture predictor transforms z into the architecture embedding e that has the same dimensionality as architecture codes. Finally, the router routes the embedding e into an architecture code a(i). We use optimal transport to evenly distribute the prompts in a training batch among the architecture codes. The architecture code a(i) = (u(i), v(i)) determines pruning the model’s width and depth. We train the prompt router’s parameters and architecture codes in an end-to-end manner using the denoising objective of the pruned model LDDPM, distillation loss between the pruned and original models Ldistill, average resource usage for the samples in the batch R, and contrastive objective Lcont, encouraging embeddings e preserving semantic similarity of the representations z.
Table of Contents
Installation
Follow these steps to set up the project:
1. Create Conda Environment
Use the provided env.yaml file:
conda env create -f env.yaml
2. Activate the Conda Environment
Activate the environment:
conda activate pdm
3. Install Project Dependencies
From the project root directory, install the dependencies:
pip install -e .Data Preparation
Prepare the data for training as mentioned in the paper. You can also adapt aptp for your own dataset with minor code modifications.
1. Download Conceptual Captions
Follow the instructions here to download Conceptual Captions. Place the data in a directory of your choice, maintaining the structure:
conceptual_captions
├── Train_GCC-training.tsv
├── Val_GCC-1.1.0-Validation.tsv
├── training
│ ├── 10007_560483514
│ └── ...
└── validation
├── 1852290_2006010568
└── ...
1.1 Remove Corrupt Images (Optional)
There are some urls in the Conceptual Captions dataset that are not valid. The download will result in some corrupt files that can't be opened. These could be removed to ensure a more efficient training.
2. Download MS-COCO 2014
2.1 Download the training and validation images
Download 2014 train and 2014 val images from the COCO website. Place them in your chosen directory.
2.2 Download the annotations
Download the 2014 train/val annotations and place them in the same directory as the images. Your directory should look like this:
coco
├── annotations
│ ├── captions_train2014.json
│ ├── captions_val2014.json
│ └── ...
└── images
├── train2014
│ ├── COCO_train2014_000000000009.jpg
│ └── ...
└── val2014
├── COCO_val2014_000000000042.jpg
└── ...
Training
Training is done in two stages: pruning the pretrained T2I model (Stable Diffusion 2.1 in this case) and fine-tuning each expert on the prompts assigned to it. Configuration files for both Conceptual Captions and MS-COCO are provided in the configs directory. You can use these configuration files to run the pruning process. Sample multi-node SLURM and PBS scripts can be found in cluster scripts.
1. Pruning
You can use the following command to run pruning:
accelerate launch scripts/aptp/prune.py \
--base_config_path path/to/configs/pruning/file.yaml \
--cache_dir /path/to/.cache/huggingface/ \
--wandb_run_name WANDB_PRUNING_RUN_NAME This creates a checkpoint directory named "wandb_run_name" in the logging directory specified in the config file.
2. Data Preparation for Fine-tuning
The pruning stages results in

