Skip to content

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

2 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Vision Transformer

A Vision Transformer (ViT) implementation for image classification using PyTorch. This project includes model training, evaluation, and inference capabilities.

Project Structure

vision-transformer/
├── vision_transformer.py      # Main script with model architecture and training pipeline
├── engine.py                  # Training engine with train() function
├── helper_functions.py        # Utility functions for plotting and visualization
├── predictions.py             # Inference and prediction functions
└── README.md                  # This file

Requirements

  • Python 3.7+
  • PyTorch
  • torchvision
  • torchinfo
  • matplotlib
  • seaborn
  • numpy

Installation

  1. Install PyTorch and dependencies:
pip install torch torchvision torchinfo matplotlib seaborn numpy
  1. Clone the MICCAI dataset (handled automatically in the script):
# The script automatically clones the dataset
subprocess.run(["git", "clone", "https://github.com/ssanya942/MICCAI-Educational-Challenge-2024.git"])

Usage

Run the main script:

python vision_transformer.py

Key Components

1. vision_transformer.py

Main script containing:

  • Vision Transformer model architecture
  • Data loading and preprocessing
  • Model training configuration
  • Inference and prediction pipeline

2. engine.py

Training engine with:

  • train() function: Handles model training loop
  • Supports batch processing on CUDA/CPU
  • Computes training/validation metrics
  • Returns training history

3. helper_functions.py

Visualization utilities:

  • plot_loss_curves(): Plots training and validation loss/accuracy curves

4. predictions.py

Inference functions:

  • pred_and_plot_image(): Makes predictions on custom images and visualizes results

Features

  • Vision Transformer Architecture: State-of-the-art transformer-based image classification
  • GPU Support: Automatic CUDA detection and GPU training
  • Data Loading: Efficient batch processing with DataLoader
  • Loss Tracking: Comprehensive metrics logging during training
  • Visualization: Training curves and prediction visualization

Training Parameters

Default training configuration:

  • Optimizer: Adam (lr=3e-3)
  • Loss Function: CrossEntropyLoss
  • Epochs: 30
  • Device: CUDA (if available) or CPU

Inference

To make predictions on custom images:

from predictions import pred_and_plot_image

custom_image_path = "/path/to/image.png"
pred_and_plot_image(model=vit,
                    class_names=class_names,
                    image_path=custom_image_path)

Notes

  • The dataset paths are configured for Google Colab (/content/ directory)
  • Modify paths accordingly for local execution
  • Requires GPU for efficient training of large models
  • Training time varies based on dataset size and hardware

License

Original model based on "An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale" (Dosovitskiy et al., 2020)

About

implementation of vision-Transformer Base

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages