Skip to content

Latest commit

 

History

18 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Improving Unsupervised Hierarchical Representation with Reinforcement Learning

Ruyi An  Yewen Li  Xu He  Pengjie Gu  Mengchen Zhao  Dong Li  Jianye Hao  Chaojie Wang  Bo An  Mingyuan Zhou
🎈 Accepted to CVPR 2024

• [pdf] •

If you find our project helpful, kindly consider ⭐ this repo. Thanks! 🖐️

📮 News

  • Jun. 2024: We will be presenting our paper at CVPR 2024.
  • May 2024: We released the codebase for this project.

🛠️ Installation

Codes and Environment

# clone this repository
git clone https://github.com/ruyianry/rep_hierarchy_rl.git
cd rep_hierarchy_rl

# create a new anaconda environment
conda create -n rephrl python=3.8 -y
conda activate rephrl

# install python dependencies
conda install -y -c pytorch pytorch torchvision torchaudio cudatoolkit=11.8
pip install -r requirements.txt
pip install --editable .

Ensure that the CUDA version used by torch corresponds to the one on the device.

Package testing

pytest -v --cov --cov-report=term tests

Please run the above check to ensure that the code works as expected on your system.

🏃Training

The below commands will train a HVAE with reinforcement learning on FashionMNIST and CIFAR-10 datasets.

The dataset will be downloaded automatically if it is not found in the data directory via torchvision.datasets.

FashionMNIST

python3 scripts/dvae_run_FashionMNIST_RLQ.py

CIFAR-10

python3 scripts/dvae_run_CIFAR10_RLQ.py

The other datasets can be trained by modifying the train_datasets parameter in the script.

🔝 Citation

If you find our work useful for your research, kindly consider citing our paper:

@inproceedings{hier_rep_rl,
  title={Improving Unsupervised Hierarchical Representation with Reinforcement Learning},
  author={An, Ruyi and Li, Yewen and He, Xu and Gu, Pengjie and Zhao, Mengchen and Li, Dong and Hao, Jianye and An, Bo and Wang, Chaojie and Zhou, Mingyuan},
  booktitle={{CVPR}},
  year={2024}
}

🖖 Acknowledgement

This implementation is based on the following repositories:

🫡 Salute!

☎️ Contact

Please feel free to reach us out at ran003😎ntu.edu.sg should you need any help.

About

[CVPR'24] Official Implementation of "Improving Unsupervised Hierarchical Representation with Reinforcement Learning"

Resources

Stars

8 stars

Watchers

1 watching

Forks

Used by

Contributors

Languages