--- 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%**