Overcoming Generic Knowledge Loss with Selective Parameter Update (CVPR'24)
July 1, 2024 · View on GitHub
Authors: Wenxuan Zhang, Paul Janson, Rahaf Aljundi, Mohamed Elhoseiny @ KAUST Vision-CAIR, TME
Use this repo to reproduce the results of our methods as well as the baselines.
Installation
conda env create -f environment.yml
conda activate clip
Use the following to install the learning rate scheduler
pip install 'git+https://github.com/katsura-jp/pytorch-cosine-annealing-with-warmup'
Reproduce Our Results
Dataset
Prepare the datasets by following the instructions in the data folder.
- CIFAR 100, FGVC-Aircraft, GTSRB: Automatically downloaded by torchvision.
- CUB-200-2011, Stanford Cars: Automatically downloaded by Huggingface.
- Birdsnap: Follow the download instructions from here.
- CC12M: Follow the download instructions from here.
- ImageNet 1k: Follow the download instructions from the official website.
Reproduce our method
python main.py dataset=[cifar100 | cub | cars | aircraft | gtsrb | birdsnap ]
Run baselines
Supported baselines:
flyp: Finetune like you pretrain [paper]er: Experimence replay [paper]lwf: Learning without forgetting [paper]mas: Memory aware synapses [paper]prd: Prototype-sample relation distillation [paper]loraewc: LoRA finetune with EWC regularization[paper]slca: Slow learner with classifier alignment [paper]sparsecl: Sparse Continual Learning [paper]spg: Soft-masks parameter updating [paper]zscl: Zero-shot Continual Learning [paper]
python main.py \
dataset=[cifar100 | cub | cars | aircraft | gtsrb | birdsnap ] \
baseline@_global_=[flyp | er | lwf | mas | prd | loraewc | slca | sparsecl | spg | zscl]
Supported features
For replay based method, use balanced_buffer=False to apply uniform sampling (uniformly from buffer and the current task)
python main.py dataset=your_dataset baseline@_global_=your_baseline balanced_buffer=False
Use joint=True for joint training
python main.py dataset=your_dataset baseline@_global_=your_baseline joint=True
Adjust buffer_size to scaling down or up the buffer size
python main.py dataset=your_dataset baseline@_global_=your_baseline buffer_size=0.5
Adjust num_tasks to adjust the number of split of dataset
python main.py dataset=your_dataset baseline@_global_=your_baseline num_tasks=20
Acknowledgement
- Learning rate scheduler: Cosine Annealing with Warmup for PyTorch
- CLIP backbone: OpenCLIP
- Classic continual learning methods: Avalanche
Citation
@inproceedings{zhang2024overcoming,
title={Overcoming Generic Knowledge Loss with Selective Parameter Update},
author={Zhang, Wenxuan and Janson, Paul and Aljundi, Rahaf and Elhoseiny, Mohamed},
booktitle={CVPR},
year={2024}
}