EuroPython2022/pyro-vision
1
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 