This repository contains PruneNet, a novel model compression framework that uses reinforcement learning to compress large language models without requiring calibration data.
Based on the paper: You Only Prune Once: Designing Calibration-Free Model Compression With Policy Learning
- 🎯 No Calibration Data Required - Learns compression policy directly from model weights
- 🤖 Reinforcement Learning-Based - Learns optimal neuron selection strategy
- 📊 Preserves Spectral Properties - Maintains weight matrix characteristics
- 🚀 Easy to Use - Simple
fit()andcompress()API following scikit-learn patterns - 🔧 Flexible Configuration - Extensive hyperparameter control
- 📦 Multiple Architectures - Supports OPT, Llama, Phi, Falcon
git clone https://github.com/parmanu-lcs2/efficient_pruners
cd efficient_pruners
pip install -e .pip install efficient-prunersfrom efficient_pruners import PruneNet, PruningConfig
# Configure hyperparameters
config = PruningConfig(
num_episodes=20,
learning_rate=0.001
)
# Initialize pruner
pruner = PruneNet(config)
# Train policy on specific model with target compression ratio
pruner.fit(model_name="facebook/opt-125m", compression_ratio=0.3)
# Compress with the same or different ratio
compressed_model = pruner.compress(compression_ratio=0.3)
# Save compressed model
compressed_model.save_pretrained("./compressed_model")
# Test text generation with compressed LLM
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained("facebook/opt-125m")
inputs = tokenizer("The future of AI is", return_tensors="pt")
# Generate text with compressed model
outputs = compressed_model.generate(**inputs, max_length=50)
text = tokenizer.decode(outputs[0], skip_special_tokens=True)
print(text)The original command-line interface is still available in the prunenet/ directory:
python3 -m prunenet \
--model_name facebook/opt-125m \
--compression_ratio 0.3 \
--save_dir ./models/ \
--device cuda:0- API Guide - Complete API reference
- Test Notebook - Interactive fit/compress test with visualizations
- Test Script - Automated fit/compress test
PruneNet/
├── src/efficient_pruners/ # Main package
│ ├── core.py # PruneNet class (fit/compress API)
│ ├── config.py # PruningConfig dataclass
│ ├── models/ # SparsityPredictor policy network
│ │ └── sparsity_predictor.py
│ └── utils/ # Model and reward utilities
│ ├── model_utils.py
│ └── reward_utils.py
├── examples/ # Test & usage examples
│ └── test_fit_compress.py # Complete fit/compress test script
├── notebooks/ # Interactive tutorials
│ └── test_fit_compress.ipynb # Complete fit/compress test notebook
├── docs/ # Documentation
│ └── API_GUIDE.md
├── prunenet/ # Original CLI implementation
├── setup.py # Package setup
├── pyproject.toml # Modern build system
└── requirements.txt # Dependencies
- OPT: facebook/opt-125m, facebook/opt-1.3b, etc.
- Llama: meta-llama/Llama-2-7b-hf, etc.
- Phi: microsoft/phi-1, microsoft/phi-2, etc.
- Falcon: tiiuae/falcon-7b, etc.
Run the comprehensive test to verify both fit() and compress() methods:
python examples/test_fit_compress.pyThis script will:
- ✅ Train an RL policy using
fit() - ✅ Compress the model using
compress() - ✅ Test
.generate()on the compressed LLM - ✅ Compare outputs between original and compressed models
- ✅ Display compression statistics
jupyter notebook notebooks/test_fit_compress.ipynbThe notebook includes:
- Step-by-step walkthrough of
fit()andcompress() - Visualizations of training progress
- Interactive text generation testing with compressed model
- Side-by-side comparison of model outputs
config = PruningConfig(
num_episodes=20,
learning_rate=0.001,
use_kld=True, # Enable KL divergence regularization
gamma=0.99, # Reward discount factor
device="auto", # Auto-detect GPU/CPU
save_dir="./outputs" # Checkpoint directory
)
pruner = PruneNet(config)
pruner.fit(model_name="facebook/opt-125m")
compressed_model = pruner.compress(compression_ratio=0.3)See API_GUIDE.md for all configuration options.
Typical compression results on OPT-125M:
| Compression | Size Reduction | Perplexity Impact |
|---|---|---|
| 20% | ~15% | +2-3% |
| 30% | ~22% | +3-5% |
| 40% | ~30% | +5-8% |
| 50% | ~37% | +8-12% |
The original research scripts are preserved in prunenet/ and experiments/ directories. See the original README sections below for research-specific details.
If you find our work useful in your projects/research, kindly cite our paper:
@inproceedings{
sengupta2025you,
title={You Only Prune Once: Designing Calibration-Free Model Compression With Policy Learning},
author={Ayan Sengupta and Siddhant Chaudhary and Tanmoy Chakraborty},
booktitle={The Thirteenth International Conference on Learning Representations},
year={2025},
url={https://openreview.net/forum?id=5RZoYIT3u6}
}