Team Ai
Apppublic

EuroPython2022/pyro-vision

sourceHugging Faceapache-2.0updated 4y agoView on Hugging Face
1likes
app.py73 linesDownload Raw Back to root
1# Copyright (C) 2022, Pyronear.2 3# This program is licensed under the Apache License 2.0.4# See LICENSE or go to <https://www.apache.org/licenses/LICENSE-2.0> for full license details.5 6import argparse7import json8 9import gradio as gr10import numpy as np11import onnxruntime12from huggingface_hub import hf_hub_download13from PIL import Image14 15REPO = "pyronear/rexnet1_0x"16 17 18# Download model config & checkpoint19with open(hf_hub_download(REPO, filename="config.json"), "rb") as f:20    cfg = json.load(f)21 22ort_session = onnxruntime.InferenceSession(hf_hub_download(REPO, filename="model.onnx"))23 24def preprocess_image(pil_img: Image.Image) -> np.ndarray:25    """Preprocess an image for inference26 27    Args:28        pil_img: a valid pillow image29 30    Returns:31        the resized and normalized image of shape (1, C, H, W)32    """33 34    # Resizing (PIL takes (W, H) order for resizing)35    img = pil_img.resize(cfg["input_shape"][-2:][::-1], Image.BILINEAR)36    # (H, W, C) --> (C, H, W)37    img = np.asarray(img).transpose((2, 0, 1)).astype(np.float32) / 25538    # Normalization39    img -= np.array(cfg["mean"])[:, None, None]40    img /= np.array(cfg["std"])[:, None, None]41 42    return img[None, ...]43 44def predict(image):45    # Preprocessing46    np_img = preprocess_image(image)47    ort_input = {ort_session.get_inputs()[0].name: np_img}48 49    # Inference50    ort_out = ort_session.run(None, ort_input)51    # Post-processing52    probs = 1 / (1 + np.exp(-ort_out[0][0]))53 54    return {class_name: float(conf) for class_name, conf in zip(cfg["classes"], probs)}55 56 57img = gr.inputs.Image(type="pil")58outputs = gr.outputs.Label(num_top_classes=1)59 60 61gr.Interface(62    fn=predict,63    inputs=[img],64    outputs=outputs,65    title="PyroVision: image classification demo",66    article=(67        "<p style='text-align: center'><a href='https://github.com/pyronear/pyro-vision'>"68        "Github Repo</a> | "69        "<a href='https://pyronear.org/pyro-vision/'>Documentation</a></p>"70    ),71    live=True,72).launch()73