This project is a reimplementation of the Zero-Shot Continual Learning (ZSCL) framework. It finetunes the CLIP model on a sequence of downstream classification tasks while preventing catastrophic forgetting of its original zero-shot transfer capabilities.
This rewrite specifically leverages Hugging Face's accelerate library for highly efficient, mixed-precision, and multi-GPU distributed training out-of-the-box.
It can run on double RTX5070Ti GPUs with batch size of 32 (using approximately 26GB VRAM), and you can also enable gradient accumulation to further increase the effective batch size.
- Standard ZSCL Continual Learning Protocol: By default, the model sequentially trains and evaluates on 11 downstream datasets in the following order:
aircraft->caltech101->cifar100->dtd->eurosat->flowers->food->mnist->oxford_pet->stanford_cars. - Accelerate Integration: Seamless support for multi-GPU training (
bf16by default), gradient clipping, and gradient accumulation. - Reference Distillation: Preserves zero-shot capabilities through logit and feature-level distillation from a frozen pre-trained teacher model, using ImageNet as the image reference dataset and Conceptual Captions as the text reference.
- Weight Ensembling & Consolidation (WE & WC): Implements moving-average weight ensembling (
merge_we) and L2 weight consolidation (l2_wc_loss) to stabilize training and mitigate forgetting across tasks.
This project uses Python 3.10+ and uv as the package manager. Dependencies are defined in pyproject.toml.
To install dependencies:
uv syncStart the training script using accelerate launch.
source .venv/bin/activate
accelerate config # configure the distributed training environment
accelerate launch train_accelerate.py \
--model "ViT-B/16" \
--datasets aircraft caltech101 cifar100 dtd eurosat flowers food mnist oxford_pet stanford_cars \
--batch_size 32 \
--iterations 1000 \
--lr 1e-5 \
--distill_temperature 2.0 \
--ce_weight 1.0 \
--l2_wc_weight 1.0 \
--logit_distill_weight 1.0--model: The CLIP model architecture (e.g.,ViT-B/16,ViT-L/14).--datasets: A list of dataset sequence names to continually train on.--output_dir: Directory to store checkpoints, logged metrics, and run history.--eval_freq: Frequency of epochs between evaluations.--freeze_visual_layers: Integer representing how many bottom layers of the vision transformer to freeze (0 to train the entire vision backbone).
train_accelerate.py: The main entry point containing the continual training loop, data loading, distillation calculation, and metric aggregation.evaluation.py: Functions to perform zero-shot classification evaluation across all testing datasets.optimizations.py: Advanced optimization functions, includingdistillation, L2 weight consolidation (l2_wc_loss), and moving-average weight ensembling (merge_we).utils.py: Contains training utilities, CLI argument parsing, model freezing/copying, and checkpoint management viaRunManager.
