Skip to content

About

[ICML'24] Code for pruning vision MoE (VMoE)

Resources

Stars

0 stars

Watchers

0 watching

Forks

Repository files navigation

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.

Installation

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_transformer

Pruning and Saving Checkpoints

Use the script below to prune experts from a finetuned VMoE checkpoint and save the pruned model.

Example Command

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

Arguments

  • --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.

Output

The script saves a pruned checkpoint at the location specified by --output, which can be used for further finetuning or inference.

Inference on Pruned VMoE Model

Use the following script to run inference using a pruned VMoE checkpoint.

Example Command

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

Arguments

  • --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).

Output

The script reports evaluation metrics (e.g., top-1 / top-5 accuracy) on the specified dataset split using the pruned model checkpoint.

Finetuning the Pruned Model

Use the following script to finetune a pruned VMoE checkpoint on a downstream dataset.

Example Command

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=500

Note: Replace /path/to/moe_pruning/ with the absolute path to the repository on your system.

Arguments

  • --workdir : Directory for logs, checkpoints, and training artifacts.
  • --pruned_model : Path to the pruned model checkpoint (without .pkl extension 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.

Output

The finetuned pruned checkpoint is saved to the path specified by --savefile and can be used for inference or further training.

Citation

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}
}

About

[ICML'24] Code for pruning vision MoE (VMoE)

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages