Deep Learning powered system for cataract classification and severity estimation
Built with PyTorch, Torchvision, and ResNet architectures.
This repository contains two complementary modules:
- main.py – Trains a
ResNet50model for severity prediction (regression output as percentage). - test.py – Loads a
ResNet18model for binary classification (Cataract vs No Cataract) and predicts on new images.
- 📊 Severity Prediction – Outputs cataract severity as a percentage using regression (MSE loss).
- 🩺 Binary Classification – Distinguishes between Cataract and No Cataract cases.
- ⚡ Transfer Learning – Fine-tunes pre-trained ResNet models for medical imaging tasks.
- 🎯 Data Augmentation – Includes resizing, rotation, color jitter, and normalization for robust training.
- 💾 Model Persistence – Saves trained weights for later inference.
- Frameworks: PyTorch, Torchvision
- Models: ResNet50 (regression), ResNet18 (classification)
- Tools: Matplotlib, PIL
python main.py
This will:
- Train
ResNet50on cataract images - Save weights to
cataract_severity_model.pth - Report mean severity prediction error
python test.py
This will:
- Load
ResNet18with trained weights (cataract_model.pth) - Predict Cataract vs No Cataract for a given image
- Print the predicted class
├── main.py # Train ResNet50 for severity regression ├── test.py # Test ResNet18 for binary classification ├── Cataract/ │ └── processed_images/ │ ├── train/ │ └── test/ └── README.md