A deep learning pipeline for multi-label classification of retinal diseases from fundus images using state-of-the-art Vision Transformer architectures and transfer learning techniques.
Author: Chidwipak Kuppani
Note
Development Context This project was developed in July 2025 on a remote SSH college system (GPU server: RTX 6000) provided for research and coursework. All development, training, and evaluation was conducted on these remote servers, which is why the project is being pushed to GitHub now (January 2026) rather than during the original development period.
Why Now? As I'm applying for internships and research positions, I'm consolidating my work from various remote systems into a public portfolio on GitHub. This project represents authentic research work completed during my academic studies, now being shared for professional opportunities.
This project implements and compares multiple pre-trained Vision Transformer models for detecting ocular diseases from the ODIR-5K dataset (Ocular Disease Intelligent Recognition). The system performs multi-label classification of fundus images into 4 categories:
| Code | Disease | Description |
|---|---|---|
| N | Normal | Healthy retina with low optic cup to disc ratio |
| D | Diabetic Retinopathy | Micro-aneurysms, hemorrhages (red spots), exudates (yellow spots) |
| C | Cataract | Blurred/absent basic anatomical structures |
| M | Pathological Myopia | Peri-papillary atrophy around optic disc |
- π₯ Multi-label classification - Detect multiple conditions per image
- π¬ 4 Model Architectures - ViT, DeiT, Swin Transformer, ResNet-50
- π Attention Visualization - Interpretability through attention maps
- β‘ PyTorch Lightning - Scalable, professional training pipeline
- π€ HuggingFace Transformers - State-of-the-art pretrained models
- π Comprehensive Metrics - F1-Score, Macro F1, Ranking Average Precision
| Model | Normal (N) | Diabetic Retinopathy (D) | Cataract (C) | Myopia (M) | F1 Macro | Ranking AP |
|---|---|---|---|---|---|---|
| Swin | 80.8% | 61.1% | 86.3% | 95.9% | 81.1% | 81.0% |
| DeiT | 78.9% | 63.5% | 85.1% | 98.0% | 81.4% | 80.8% |
| ViT | 76.8% | 63.4% | 83.3% | 95.9% | 79.9% | 78.6% |
| ResNet-50 | 80.1% | 53.7% | 79.1% | 98.0% | 77.7% | 79.4% |
Key Finding: Vision Transformers consistently outperform the CNN baseline (ResNet-50) across all metrics.
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β F1 Score Performance (%) β
ββββββββββββββββββββββββ¬βββββββββ¬βββββββββ¬βββββββββ¬ββββββββββββ¬βββββββββ€
β Class β Swin β DeiT β ViT β ResNet-50 β Best β
ββββββββββββββββββββββββΌβββββββββΌβββββββββΌβββββββββΌββββββββββββΌβββββββββ€
β Normal (N) β 80.8 β 78.9 β 76.8 β 80.1 β Swin β
β Diabetic Retinop.(D) β 61.1 β 63.5 β 63.4 β 53.7 β DeiT β
β Cataract (C) β 86.3 β 85.1 β 83.3 β 79.1 β Swin β
β Myopia (M) β 95.9 β 98.0 β 95.9 β 98.0 β DeiT β
ββββββββββββββββββββββββ΄βββββββββ΄βββββββββ΄βββββββββ΄ββββββββββββ΄βββββββββ
- Swin Transformer achieves best overall performance - Highest Ranking Average Precision (81.0%)
- Myopia detection is highly accurate - All models achieve >95% F1
- Diabetic Retinopathy is most challenging - Best F1 is 63.5% (DeiT)
- Vision Transformers outperform CNN baseline - Swin exceeds ResNet by 3.4% in F1 Macro
| Label | Class | Training | Validation | Testing | Total |
|---|---|---|---|---|---|
| N | Normal | 1337 | 325 | 402 | 2064 |
| D | Diabetic Retinopathy | 1079 | 285 | 354 | 1718 |
| C | Cataract | 195 | 49 | 58 | 302 |
| M | Pathological Myopia | 163 | 39 | 53 | 255 |
| Parameter | Value |
|---|---|
| Optimizer | AdamW |
| Learning Rate | 5Γ10β»β΄ |
| Batch Size | 16 |
| Max Epochs | 400 |
| Early Stopping | 10 epochs patience |
| Loss Function | Weighted BCE (class-balanced) |
| GPU | RTX 6000 |
OcularAI/
βββ core/ # Core modules
β βββ models/
β β βββ classifiers.py # 4 model implementations
β βββ data/
β β βββ dataset.py # Dataset class
β βββ config.py # Hyperparameters
β βββ utils.py # Utilities
βββ pipeline/ # Training pipeline
β βββ train.py # Training script
β βββ visualization/
β β βββ attention.py # Attention maps
β βββ analysis/
β βββ eda.py # Exploratory analysis
β βββ statistics.py # Dataset statistics
βββ outputs/ # Training outputs
β βββ training_logs/ # Experiment logs
βββ data/ # Dataset files
β βββ ODIR/
βββ requirements.txt
βββ README.md
# Clone the repository
git clone https://github.com/YOUR_USERNAME/OcularAI.git
cd OcularAI
# Install dependencies
pip install -r requirements.txtcd pipeline
python train.pyConfigure the model in core/config.py:
# Options: SWIN, VIT, DeiT, ResNet
model_processor = 'SWIN' # Best performing modelcd pipeline/visualization
python attention.py| Model | Pre-trained Source | Key Features |
|---|---|---|
| Swin | microsoft/swin-tiny-patch4-window7-224 | Shifted window attention, hierarchical design |
| DeiT | facebook/deit-base-distilled-patch16-224 | Knowledge distillation, data-efficient |
| ViT | google/vit-base-patch16-224 | Original vision transformer architecture |
| ResNet-50 | microsoft/resnet-50 | CNN baseline with residual connections |
ODIR-5K (Ocular Disease Intelligent Recognition)
- Source: Peking University / Shanggong Medical Technology Co., Ltd.
- Images: ~5000 pairs of left/right fundus photographs
- Split: Train (70%) / Validation (15%) / Test (15%)
- Preprocessing: Quality filtering, 224Γ224 resize, stratified splitting
- Python 3.10+
- PyTorch 2.2+
- PyTorch Lightning 2.2+
- HuggingFace Transformers 4.39+
- CUDA-capable GPU (recommended: RTX 6000 or equivalent)
If you use this work, please cite:
@software{ocularai2024,
author = {Kuppani, Chidwipak},
title = {OcularAI: Multi-Label Retinal Disease Classification using Vision Transformers},
year = {2024},
url = {https://github.com/YOUR_USERNAME/OcularAI}
}- ODIR-5K Dataset - Peking University
- HuggingFace Transformers
- PyTorch Lightning
MIT License