Aditya-Sai-19/diabetic-retinopathy-swin
025
Diabetic Retinopathy Prediction — MobileNetV2
This model predicts the severity grade of diabetic retinopathy (DR) from retinal fundus images. It classifies images into 5 grades:
Model Details
- Architecture: MobileNetV2 (2.2M parameters) fine-tuned for 5-class DR grading
- Input: 224×224 RGB retinal fundus images
- Training Data: youssefedweqd/Diabetic_Retinopathy_Detection — 25,290 train / 2,810 val / 7,026 test images
- Loss Function: Focal Loss (γ=2.0) with inverse-frequency class weights — critical for handling severe class imbalance
- Primary Metric: Quadratic Weighted Kappa (QWK) — the standard metric for ordinal DR grading
Training Recipe
Based on published DR detection literature (see References below):
Class Distribution (severe imbalance)
Results
Validation Performance (best epoch = 3)
Test Set Performance
- Accuracy: 13.5%
- QWK: 0.384
Per-Class Test Results
Usage
from transformers import pipeline
classifier = pipeline("image-classification", model="Aditya-Sai-19/diabetic-retinopathy-swin")
result = classifier("path/to/retinal_image.jpg")
print(result)
# [{'label': 'Moderate', 'score': 0.45}, {'label': 'Mild', 'score': 0.22}, ...]Or manually:
from transformers import AutoImageProcessor, AutoModelForImageClassification
from PIL import Image
import torch
processor = AutoImageProcessor.from_pretrained("Aditya-Sai-19/diabetic-retinopathy-swin")
model = AutoModelForImageClassification.from_pretrained("Aditya-Sai-19/diabetic-retinopathy-swin")
image = Image.open("retinal_scan.jpg")
inputs = processor(image, return_tensors="pt")
with torch.no_grad():
logits = model(**inputs).logits
predicted_class = logits.argmax(-1).item()
print(f"Predicted: {model.config.id2label[predicted_class]}")Limitations & Recommendations for Improvement
This model was trained on CPU only (free tier) with a lightweight MobileNetV2 backbone. For production-quality DR grading, consider:
- Use a larger backbone: Swin Transformer Base or EfficientNet-B4 (published SOTA: QWK > 0.93 on APTOS)
- Train on GPU: Enables larger batch sizes, more epochs, and larger models
- Ben Graham preprocessing: Apply contrast enhancement (
I_ce = 4*I - 4*G(σ=w/90)*I + 128) and FOV cropping — this alone adds +2-5% AUC - More data: Merge APTOS + Messidor-2 + IDRiD + EyePACS for better generalization
- Data augmentation: 10x oversampling of minority classes (grades 3 & 4)
- Post-processing: Optimized ordinal thresholds for QWK maximization
References
- Dual Branch Deep Learning Network for Detection and Stage Grading of DR (arXiv:2308.09945) — QWK=0.930
- A ResNet is All You Need? (arXiv:2210.03180) — AUC=0.955
- Ord2Seq: Ordinal Regression as Label Sequence Prediction (arXiv:2307.09004) — Acc=84.2%
Framework Versions
- Transformers 5.6.2
- PyTorch 2.11.0
- Datasets 4.8.4
