Skip to content
witheredAdPublic

About

A reimplementation of the Zero-Shot Continual Learning (ZSCL) framework, using Hugging Faces's accelerate library.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Repository files navigation

ZSCL (Zero-Shot Continual Learning) - Accelerate Rewrite

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.

Features

  • 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 (bf16 by 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.

Requirements

This project uses Python 3.10+ and uv as the package manager. Dependencies are defined in pyproject.toml.

To install dependencies:

uv sync

Usage

Start 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

Key Arguments

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

Code Structure

  • 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, including distillation, 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 via RunManager.

Results

Metrics

About

A reimplementation of the Zero-Shot Continual Learning (ZSCL) framework, using Hugging Faces's accelerate library.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages