nailarais1's picture
Update README.md
e1ca410 verified
|
Raw
History Blame Contribute Delete
5.8 kB
---
library_name: pytorch
tags:
- pytorch
- image-classification
- flowers
- computer-vision
pipeline_tag: image-classification
language:
- en
---
# 🌸 102-Flower Image Classifier — EfficientNet-B0
A PyTorch EfficientNet-B0 image classification model trained to recognize **102 flower categories** from the Oxford 102 Category Flower Dataset.
## Model Performance
| Metric | Result |
| ------------------------ | --------------- |
| Architecture | EfficientNet-B0 |
| Number of classes | 102 |
| Input size | 224 × 224 |
| Best validation accuracy | **94.38%** |
| Training epochs | 3 |
| Optimizer | AdamW |
| Learning rate | 0.001 |
The model was trained using transfer learning with an ImageNet-pretrained EfficientNet-B0 backbone.
## Dataset
This model was trained using the **Oxford 102 Category Flower Dataset**, created by **Maria-Elena Nilsback and Andrew Zisserman**.
The dataset contains 102 flower categories with variations in scale, pose, lighting, and appearance.
Official dataset page:
https://www.robots.ox.ac.uk/~vgg/data/flowers/102/
Please review the original dataset documentation and terms before using or redistributing dataset-derived material.
## Files
* `checkpoint.pth` — trained PyTorch checkpoint
* `model_config.json` — model architecture information
* `training_config.json` — training configuration
* `class_config.json` — exact class/index mappings
* `labels.txt` — flower labels
* `requirements.txt` — Python dependencies
## Checkpoint Contents
The `checkpoint.pth` file contains:
* `epoch`
* `model_state_dict`
* `optimizer_state_dict`
* `class_to_idx`
## Use the Model
Install the dependencies:
```bash
pip install torch torchvision pillow
```
Load the model:
```python
import json
import torch
import torch.nn as nn
from torchvision import models, transforms
from PIL import Image
checkpoint = torch.load(
"checkpoint.pth",
map_location="cpu",
weights_only=False
)
model = models.efficientnet_b0(weights=None)
model.classifier[1] = nn.Linear(
model.classifier[1].in_features,
102
)
model.load_state_dict(
checkpoint["model_state_dict"]
)
model.eval()
with open(
"class_config.json",
"r",
encoding="utf-8"
) as f:
class_config = json.load(f)
idx_to_class = {
int(k): v
for k, v in class_config["idx_to_class"].items()
}
transform = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(
[0.485, 0.456, 0.406],
[0.229, 0.224, 0.225]
)
])
image = Image.open(
"flower.jpg"
).convert("RGB")
x = transform(
image
).unsqueeze(0)
with torch.inference_mode():
probabilities = torch.softmax(
model(x),
dim=1
)
confidence, prediction = probabilities.max(
dim=1
)
idx = prediction.item()
print(
"Prediction:",
idx_to_class[idx]
)
print(
"Confidence:",
f"{confidence.item() * 100:.2f}%"
)
```
## Top-5 Predictions
You can also get the five most likely flower categories:
```python
with torch.inference_mode():
probabilities = torch.softmax(
model(x),
dim=1
)
values, indices = torch.topk(
probabilities,
k=5
)
for probability, index in zip(
values[0],
indices[0]
):
flower = idx_to_class[index.item()]
confidence = probability.item() * 100
print(
f"{flower}: {confidence:.2f}%"
)
```
## Training Configuration
The model was trained using transfer learning.
* Architecture: EfficientNet-B0
* Classes: 102
* Image size: 224 × 224
* Batch size: 32
* Epochs: 3
* Optimizer: AdamW
* Learning rate: 0.001
* Loss: CrossEntropyLoss
* Scheduler: StepLR
* Mixed precision: CUDA when available
### Training Augmentation
* Random resized crop
* Random horizontal flip
* Color jitter
* ImageNet normalization
### Validation Preprocessing
* Resize to 256
* Center crop to 224
* ImageNet normalization
## Evaluation
The best validation accuracy achieved during training was:
**94.38%**
This result corresponds to the validation split used during training.
Performance may vary on images that differ substantially from the training data.
## Interactive Demo
An interactive Gradio application can be deployed using this model so that users can upload flower images directly through a web browser.
The demo can provide:
* Image upload
* Flower prediction
* Confidence score
* Top-5 predictions
## Limitations
This model is designed to classify images into the 102 flower categories represented in the training dataset.
Predictions may be less reliable when:
* The image does not contain a supported flower category.
* The flower is heavily obscured.
* The image is blurry or poorly illuminated.
* Multiple flowers appear in the image.
* The image differs substantially from the training distribution.
This model should be considered an image-classification research/demo model and not a definitive botanical identification system.
## Citation
If you use this model or the underlying dataset, please provide attribution to the original dataset authors.
**Maria-Elena Nilsback and Andrew Zisserman**
*"Automated Flower Classification over a Large Number of Classes."*
Proceedings of the Indian Conference on Computer Vision, Graphics and Image Processing (ICVGIP), 2008.
## Dataset Reference
Oxford 102 Category Flower Dataset:
https://www.robots.ox.ac.uk/~vgg/data/flowers/102/
## Author
**Naila Rais**
Hugging Face:
`nailarais1`
Model:
`nailarais1/image-classifier-efficientnet`
Architecture:
**EfficientNet-B0**
Best validation accuracy:
**94.38%**