| import json |
| from functools import lru_cache |
| from urllib.parse import quote_plus |
|
|
| import gradio as gr |
| import torch |
| import torch.nn.functional as F |
| from huggingface_hub import hf_hub_download |
| from PIL import Image |
| from torchvision import models, transforms |
|
|
| |
| MOBILENET_REPO = "cpoisson/plantnet300k-mobilenetv3-small" |
| RESNET_REPO = "cpoisson/plantnet300k-resnet18" |
| NUM_CLASSES = 1081 |
|
|
| MODEL_CHOICES = { |
| "MobileNetV3-Small (v2) β 10 MB Β· 3.9M params Β· +1.56% accuracy β": "mobilenet_v2", |
| "MobileNetV3-Small (v1) β 10 MB Β· 3.9M params Β· baseline": "mobilenet", |
| "ResNet18 β 45 MB Β· 11.7M params Β· reference model": "resnet18", |
| } |
|
|
| |
| _json_path = hf_hub_download(MOBILENET_REPO, "plantnet300K_species_id_2_name.json") |
| with open(_json_path) as f: |
| _id2name = json.load(f) |
| class_ids = sorted(int(k) for k in _id2name) |
| class_names = [_id2name[str(cid)] for cid in class_ids] |
|
|
| |
| transform = transforms.Compose([ |
| transforms.Resize(256), |
| transforms.CenterCrop(224), |
| transforms.ToTensor(), |
| transforms.Normalize(mean=[0.485, 0.456, 0.406], |
| std=[0.229, 0.224, 0.225]), |
| ]) |
|
|
| |
| @lru_cache(maxsize=3) |
| def load_model(key: str) -> torch.nn.Module: |
| if key == "mobilenet_v2": |
| m = models.mobilenet_v3_small(weights=None, num_classes=NUM_CLASSES) |
| path = hf_hub_download(MOBILENET_REPO, "mobilenetv3_small_v2.pth") |
| elif key == "mobilenet": |
| m = models.mobilenet_v3_small(weights=None, num_classes=NUM_CLASSES) |
| path = hf_hub_download(MOBILENET_REPO, "plantnet_mobilenetv3.pth") |
| else: |
| m = models.resnet18(weights=None, num_classes=NUM_CLASSES) |
| path = hf_hub_download(RESNET_REPO, "plantnet_resnet18.pth") |
| m.load_state_dict(torch.load(path, map_location="cpu", weights_only=True)) |
| return m.eval() |
|
|
| |
| def classify(image_path: str, model_label: str, top_k: int): |
| if image_path is None: |
| return {}, "" |
| key = MODEL_CHOICES[model_label] |
| model = load_model(key) |
| img = Image.open(image_path).convert("RGB") |
| tensor = transform(img).unsqueeze(0) |
| with torch.no_grad(): |
| probs = F.softmax(model(tensor), dim=1)[0] |
| topk_probs, topk_idx = probs.topk(top_k) |
| predictions = { |
| class_names[i.item()]: float(p) |
| for i, p in zip(topk_idx, topk_probs) |
| } |
| top_name = next(iter(predictions)) |
| search_url = ( |
| "https://www.inaturalist.org/taxa/search?q=" |
| + quote_plus(" ".join(top_name.split()[:2])) |
| ) |
| link_html = ( |
| f'<a href="{search_url}" target="_blank" style="' |
| f'display:inline-flex;align-items:center;gap:8px;margin-top:10px;' |
| f'padding:10px 20px;background:#4a7c59;color:#fff;border-radius:8px;' |
| f'text-decoration:none;font-weight:600;font-size:15px;">' |
| f'π Search <em style="margin:0 4px">{top_name}</em> on iNaturalist</a>' |
| ) |
| return predictions, link_html |
|
|
| |
| ABOUT = """ |
| ## π§ͺ Experiment: small local models for plant identification |
| |
| **Core question** β *How far can a sub-15 MB model go on a real-world, |
| fine-grained botanical dataset?* Can it be useful enough to run entirely |
| offline on a phone or embedded device? |
| |
| This Space presents two fine-tuned models trained on |
| [Pl@ntNet-300K](https://zenodo.org/records/5645731) and evaluated on its |
| held-out test set. |
| |
| > β‘ **Want to run fully offline in your browser β no server, no internet after first load?** |
| > Try the [offline demo β](https://huggingface.co/spaces/cpoisson/plantnet300k-offline) |
| > Built with ONNX Runtime Web + React, models load from HF Hub and run entirely client-side. |
| |
| --- |
| |
| ### Dataset β Pl@ntNet-300K |
| |
| | | | |
| |---|---| |
| | Source | [Zenodo β DOI:10.5281/zenodo.5645731](https://zenodo.org/records/5645731) | |
| | Paper | Garcin et al., *NeurIPS 2021 Datasets & Benchmarks* | |
| | Images | 306,146 | |
| | Species | **1,081** | |
| | Train split | 245,402 images | |
| | Val split | 29,892 images | |
| | Test split | 31,112 images | |
| | Key challenge | Long-tailed: 80% of species = only 11% of images. High label ambiguity (visually similar species). | |
| |
| --- |
| |
| ### Models β Latest Results |
| |
| | Model | Version | Params | Size | Top-1 (test) | Top-5 (test) | Improvement | |
| |---|---|---|---|---|---|---| |
| | **MobileNetV3-Small** | **v2** β | 3.9M | **10 MB** | **75.45%** | **93.81%** | +1.56 pp | |
| | MobileNetV3-Small | v1 | 3.9M | 10 MB | 73.89% | 91.86% | baseline | |
| | ResNet18 | v1 | 11.7M | 45 MB | 75.82% | 93.98% | reference | |
| |
| **v2 improvements**: MobileNetV3-Small now surpasses v1 by 1.56 pp through improved training (cosine annealing, class-weighted sampling, TrivialAugment, label smoothing). Still runs at **only 10 MB** β perfect for edge deployment. |
| |
| --- |
| |
| ### Training Details |
| |
| #### v2 (Current) |
| | Parameter | Value | |
| |---|---| |
| | Optimizer | SGD (momentum=0.9, Nesterov=True) | |
| | LR Schedule | Cosine annealing (0.01 β 1e-5) | |
| | Augmentation | TrivialAugment + RandomErasing | |
| | Class Balancing | WeightedRandomSampler | |
| | Label Smoothing | 0.1 | |
| | Epochs | 60 (Phase 1: 5 frozen, Phase 2: 55 full) | |
| | Batch size | 256 | |
| | Best checkpoint | Epoch 59 (val: 75.56%, test: 75.45%) | |
| |
| #### v1 (Previous) |
| | Parameter | Value | |
| |---|---| |
| | Optimizer | Adam, lr = 1e-3 (constant) | |
| | Epochs | 60 | |
| | Batch size | 64 | |
| | Train augmentation | Resize(256) β RandomResizedCrop(224) β HFlip β ColorJitter | |
| | Loss | CrossEntropyLoss | |
| | Checkpoint | Last epoch (no best-val selection) | |
| |
| --- |
| |
| ### Known limitations (v1) |
| |
| - Adam at lr=1e-3 is aggressive for fine-tuning |
| - Weights saved at last epoch, not best checkpoint |
| - Class imbalance not addressed |
| - No LR schedule |
| |
| **v2 addressed all of these**, resulting in **+1.56 pp improvement**. |
| |
| --- |
| |
| ### Replicate |
| |
| ```bash |
| # 1. Download dataset (Zenodo) |
| wget https://zenodo.org/records/5645731/files/plantnet_300K_images.tar.gz |
| tar -xzf plantnet_300K_images.tar.gz |
| |
| # 2. Install |
| pip install torch torchvision |
| |
| # 3. Train (edit DATA_DIR at top of script) |
| python train.py # training script in each model repo |
| ``` |
| |
| Model repos: |
| [cpoisson/plantnet300k-mobilenetv3-small](https://huggingface.co/cpoisson/plantnet300k-mobilenetv3-small) Β· |
| [cpoisson/plantnet300k-resnet18](https://huggingface.co/cpoisson/plantnet300k-resnet18) |
| """ |
|
|
| |
| sample_images = [[str(p)] for p in sorted(__import__("pathlib").Path("examples").glob("*.jpg"))] |
|
|
| with gr.Blocks(title="PlantNet-300K β Small Model Experiment") as demo: |
|
|
| gr.Markdown(""" |
| # πΏ PlantNet-300K β Small Model Experiment |
| Fine-tuned **MobileNetV3-Small (10 MB)** and **ResNet18 (45 MB)** on 1,081 plant species. |
| **v2 now live: 75.45% top-1 accuracy (+1.56 pp improvement)** β |
| | β‘ [Run fully offline in your browser β](https://huggingface.co/spaces/cpoisson/plantnet300k-offline) | |
| |---|| |
| """) |
|
|
| with gr.Tabs(): |
|
|
| |
| with gr.Tab("π Classify"): |
| with gr.Row(): |
| with gr.Column(): |
| image_in = gr.Image( |
| label="Plant photo", |
| type="filepath", |
| sources=["upload", "webcam", "clipboard"], |
| ) |
| model_picker = gr.Radio( |
| choices=list(MODEL_CHOICES.keys()), |
| value=list(MODEL_CHOICES.keys())[0], |
| label="Model", |
| ) |
| top_k = gr.Slider( |
| minimum=1, maximum=10, value=5, step=1, |
| label="Top-K predictions", |
| ) |
| run_btn = gr.Button("Identify π±", variant="primary") |
|
|
| with gr.Column(): |
| label_out = gr.Label(label="Predicted species", num_top_classes=10) |
| link_out = gr.HTML() |
|
|
| gr.Examples( |
| examples=sample_images, |
| inputs=image_in, |
| label="Sample images", |
| ) |
|
|
| run_btn.click( |
| fn=classify, |
| inputs=[image_in, model_picker, top_k], |
| outputs=[label_out, link_out], |
| ) |
| image_in.change( |
| fn=classify, |
| inputs=[image_in, model_picker, top_k], |
| outputs=[label_out, link_out], |
| ) |
|
|
| |
| with gr.Tab("π About this experiment"): |
| gr.Markdown(ABOUT) |
|
|
| demo.launch() |
|
|