Code for the paper titled "A Provably Effective Method for Pruning Experts in Fine-tuned Sparse Mixture-of-Experts" [ICML'2024]
This repository implements expert-pruning, finetuning, and inference of Google's Vision Mixture-of-Experts (VMoE) model on benchmark vision tasks.
Follow the steps below to set up the environment and install dependencies:
conda create --name moe_pruning
conda activate moe_pruning
git clone https://github.com/nowazrabbani/moe_pruning.git
cd moe_pruning
git clone https://github.com/google-research/vmoe.git vision_moe
git clone https://github.com/google-research/vision_transformer.git
mv vision_moe/vmoe .
mv vision_transformer/vit_jax .
cd vision_moe
pip install -r requirements.txt
cd ..
cd vit_jax
pip install -r requirements.txt
cd ..
pip install -q 'jax[cuda]' -f https://storage.googleapis.com/jax-releases/jax_releases.html
rm -rf vision_moe
rm -rf vision_transformerUse the script below to prune experts from a finetuned VMoE checkpoint and save the pruned model.
python prune_and_save_checkpoint.py \
--finetuned_ckpt gs://vmoe_checkpoints/vmoe_b16_imagenet21k_randaug_strong_ft_ilsvrc2012 \
--pretrained_ckpt gs://vmoe_checkpoints/vmoe_b16_imagenet21k_randaug_strong \
--output pruned_ckpts/vmoe_ft_imagenet1k_router_norm_change_encoders_1357911_pruned_2_experts.pkl \
--num_experts_per_layer 8 \
--num_experts_to_prune_per_layer 2 \
--pruning_method router_norm_change \
--moe_layers_to_prune 1,3,5,7,9,11--finetuned_ckpt: Path to the finetuned checkpoint.--pretrained_ckpt: Path to the pretrained checkpoint (used as reference for pruning).--output: File path to save the pruned checkpoint.--num_experts_per_layer: Total number of experts in each MoE layer.--num_experts_to_prune_per_layer: Number of experts to prune per layer.--pruning_method: Criterion used for pruning (e.g.,router_norm_change).--moe_layers_to_prune: Comma-separated list of MoE encoder layers to prune.
The script saves a pruned checkpoint at the location specified by --output, which can be used for further finetuning or inference.
Use the following script to run inference using a pruned VMoE checkpoint.
python inference_on_pruned_vmoe_model.py \
--dataset imagenet2012 \
--split test \
--num_classes 1000 \
--checkpoint pruned_ckpts/vmoe_ft_imagenet1k_router_norm_change_encoders_1357911_pruned_2_experts.pkl \
--capacity_factor 1.5 \
--batch_size 128 \
--image_size 384 \
--patch_size 16 \
--unpruned_experts encoderblock_1=6,encoderblock_3=6,encoderblock_5=6,encoderblock_7=6,encoderblock_9=6,encoderblock_11=6--dataset: Dataset name (e.g.,imagenet2012).--split: Dataset split for evaluation (train/val/test).--num_classes: Number of output classes.--checkpoint: Path to the pruned checkpoint.--capacity_factor: MoE routing capacity factor.--batch_size: Batch size for inference.--image_size: Input image resolution.--patch_size: Patch size used in the model.--unpruned_experts: Number of remaining experts per pruned MoE layer (layer-wise specification).
The script reports evaluation metrics (e.g., top-1 / top-5 accuracy) on the specified dataset split using the pruned model checkpoint.
Use the following script to finetune a pruned VMoE checkpoint on a downstream dataset.
python finetune_pruned_model.py \
--workdir=/path/to/moe_pruning/temp \
--pruned_model=pruned_ckpts/vmoe_ft_imagenet1k_router_norm_change_encoders_1357911_pruned_2_experts \
--savefile=pruned_ckpts/vmoe_ft_imagenet1k_router_norm_change_encoders_1357911_pruned_2_experts_finetuned.pkl \
--dataset_name=imagenet2012 \
--batch_size=32 \
--num_classes=1000 \
--image_size=384 \
--evaluate_evry_steps=100 \
--train_steps=1000 \
--unpruned_experts_per_encoder=encoderblock_1=6,encoderblock_3=6,encoderblock_5=6,encoderblock_7=6,encoderblock_9=6,encoderblock_11=6 \
--lr_peak=0.00009375 \
--lr_end=1e-5 \
--lr_warmup_steps=500Note: Replace
/path/to/moe_pruning/with the absolute path to the repository on your system.
--workdir: Directory for logs, checkpoints, and training artifacts.--pruned_model: Path to the pruned model checkpoint (without.pklextension if applicable).--savefile: Output path to save the finetuned checkpoint.--dataset_name: Dataset used for finetuning (e.g.,imagenet2012).--batch_size: Training batch size.--num_classes: Number of output classes.--image_size: Input image resolution.--evaluate_evry_steps: Evaluation frequency during training.--train_steps: Total finetuning steps.--unpruned_experts_per_encoder: Number of remaining experts per pruned MoE layer.--lr_peak: Peak learning rate.--lr_end: Final learning rate.--lr_warmup_steps: Number of warmup steps.
The finetuned pruned checkpoint is saved to the path specified by --savefile and can be used for inference or further training.
If you find this repository useful, please cite the following paper:
@inproceedings{chowdhury2024provably,
title={A provably effective method for pruning experts in fine-tuned sparse mixture-of-experts},
author={Chowdhury, Mohammed Nowaz Rabbani and Wang, Meng and El Maghraoui, Kaoutar and Wang, Naigang and Chen, Pin-Yu and Carothers, Christopher},
booktitle={Proceedings of the 41st International Conference on Machine Learning},
pages={8815--8847},
year={2024}
}