PatchRot

December 11, 2024 · View on GitHub

Official Implementation of paper PatchRot: Self-Supervised Training of Vision Transformers by Rotation Prediction.

Table of Contents

Introduction

Short Summary

PatchRot is a self-supervised learning technique designed for Vision Transformers. It leverages image and patch rotation tasks to train networks to predict rotation angles, learning both global and patch-level representations.

Overview

  • PatchRot introduces a novel self-supervised strategy to learn rich and transferrable features.
  • Rotates images and image patches by 0°, 90°, 180°, or 270°.
  • Trains the network to predict the rotation angles of images and image patches as a classification task.
  • Incorporates a buffer between patches to prevent trivial solutions such as edge continuity.
  • Employs pretraining at smaller resolutions, followed by finetuning at the original size.
  • This approach encourages the model to learn both global and patch-level representations.
  • PatchRot was evaluated using the DeiT-Tiny Transformer, with dataset-specific modifications to patch size.

Usage

Requirements

Python (>= 3.8), scikit-learn, PyTorch (>= 1.10), torchvision, and timm (for defining Vision Transformers, can be replaced with other frameworks)

Run commands:

To pre-train and finetune models using PatchRot, run the following commands. Examples are provided for CIFAR10 (also available in run_cifar10.sh).

  • PatchRot Pretraining
python main_pretrain.py --dataset cifar10
  • Finetuning Pretrained Model
python main_finetune.py --dataset cifar10 --init patchrot
  • To train a baseline (Without PatchRot) set init to none: python main_finetune.py --dataset cifar10 --init none

We used a DeiT-Tiny Transformer and modified the patch size based on the dataset (refer config folder).

Data

  • To change the dataset, replace CIFAR10 with the appropriate dataset.
  • CIFAR10, CIFAR100, FashionMNIST, and SVHN are automatically downloaded by the script.
  • TinyImageNet, Animals10n, and Imagenet100 need to be downloaded manually (links below).
  • For manually downloaded datasets, use the --data_path argument to specify the path to the dataset. Example:
python main_pretrain.py --dataset tinyimagenet --data_path /path/to/data

Results

PatchRot significantly improves performance across diverse datasets. The table below compares the classification accuracy of baseline training (without PatchRot pretraining) and training with PatchRot pretraining:

DatasetWithout PatchRot PretrainingWith PatchRot Pretraining
CIFAR1084.4%91.3%
CIFAR10056.5%66.7%
FashionMNIST93.4%94.6%
SVHN92.9%96.4%
Animals10N69.6%79.5%
TinyImageNet38.4%48.8%
ImageNet10064.6%75.4%

Cite

If you found our work/code helpful, please cite our paper:

@inproceedings{Chhabra_2024_BMVC,
author    = {Sachin Chhabra and Hemanth Venkateswara and Baoxin Li},
title     = {PatchRot: Self-Supervised Training of Vision Transformers by Rotation Prediction},
booktitle = {35th British Machine Vision Conference 2024, {BMVC} 2024, Glasgow, UK, November 25-28, 2024},
publisher = {BMVA},
year      = {2024},
url       = {https://papers.bmvc2024.org/0391.pdf}
}