Create README.md
π¦· Oral Diseases Image Classification
A ResNet50-based deep learning model fine-tuned to classify six common oral diseases from intraoral images. This repository contains the best-performing model from a benchmark of four convolutional neural network architectures trained and evaluated under identical conditions.
π Best Model: ResNet50
β Accuracy: 94.77%
π― Macro F1-Score: 0.9411
π§ Framework: PyTorch
Model Overview
The model classifies the following six oral conditions:
- Calculus
- Caries
- Gingivitis
- Ulcers
- Tooth Discoloration
- Hypodontia
The final model was obtained using transfer learning with an ImageNet-pretrained ResNet50 and fine-tuned using a two-stage training strategy.
Benchmark Results
| Rank | Model | Trainable Parameters | Accuracy | Macro F1 |
|---|---|---|---|---|
| π₯ | ResNet50 | 23,520,326 | 94.77% | 0.9411 |
| π₯ | DenseNet121 | 6,960,006 | 94.51% | 0.9351 |
| π₯ | EfficientNet-B0 | 4,015,234 | 94.17% | 0.9335 |
| 4 | Scratch CNN | 11,179,590 | 83.45% | 0.8236 |
Repository Structure
checkpoints/
βββ best_model.pth
notebooks/
βββ oral-disseases-image-classification.ipynb
outputs/
βββ models_comparison.csv
βββ resnet50_confusion_matrix.png
βββ resnet50_history.png
βββ densenet121_confusion_matrix.png
βββ densenet121_history.png
βββ efficientnet_b0_confusion_matrix.png
βββ efficientnet_b0_history.png
βββ scratch_cnn_confusion_matrix.png
βββ scratch_cnn_history.png
Gradio.py
README.md
Download
Model Weights
The trained checkpoint is available in:
checkpoints/best_model.pth
or can be downloaded directly from this repository.
Dataset
Training dataset:
https://huggingface.co/datasets/nsr51324/Oral_Diseases
Original source:
Oral Diseases Dataset (Kaggle)
How to Load the Model
from huggingface_hub import hf_hub_download
import torch
weights_path = hf_hub_download(
repo_id="nsr51324/Oral_Diseases_Image_Classification",
filename="checkpoints/best_model.pth"
)
checkpoint = torch.load(weights_path, map_location="cpu")
class_names = checkpoint["class_names"]
Inference
import torch
import torch.nn as nn
from torchvision.models import resnet50
from torchvision import transforms
from PIL import Image
model = resnet50(weights=None)
model.fc = nn.Sequential(
nn.Dropout(0.3),
nn.Linear(model.fc.in_features, len(class_names))
)
model.load_state_dict(checkpoint["state_dict"])
model.eval()
transform = transforms.Compose([
transforms.Resize((224,224)),
transforms.ToTensor(),
transforms.Normalize(
[0.485,0.456,0.406],
[0.229,0.224,0.225]
)
])
image = Image.open("sample.jpg").convert("RGB")
tensor = transform(image).unsqueeze(0)
with torch.no_grad():
probabilities = torch.softmax(model(tensor), dim=1)[0]
prediction = class_names[probabilities.argmax().item()]
print(prediction)
Interactive Demo
A standalone Gradio application is included.
Run:
pip install torch torchvision gradio pillow huggingface_hub
python Gradio.py
Training Details
| Item | Value |
|---|---|
| Image Size | 224 Γ 224 |
| Batch Size | 32 |
| Epochs | Up to 30 |
| Optimizer | Adam |
| Early Stopping | Yes |
| Weight Decay | 1e-4 |
| Label Smoothing | 0.1 |
| Dropout | 0.4 |
Training consisted of two stages:
- Freeze the ResNet50 backbone and train the classifier head.
- Unfreeze the backbone and fine-tune the entire network.
Data Augmentation
The following augmentations were applied during training:
- Random Resized Crop
- Horizontal Flip
- Rotation
- Color Jitter
- Random Erasing
Evaluation
The repository includes:
- Confusion matrices
- Training history
- Classification metrics
- Model comparison
- CSV benchmark results
See the outputs/ directory for complete evaluation results.
Intended Use
This model is intended for research, educational purposes, and AI experimentation.
It is not a certified medical device and must not be used as a substitute for professional clinical diagnosis.
License
This project is released under the MIT License.
Please refer to the dataset license before commercial use.
Author
Nasr Mohamed
AI Engineer