CoolFace
Apppublic

Detomo/Image-Classification

sourceHugging Faceupdated 5y agoView on Hugging Face
5likes
app.py82 linesDownload Raw Back to root
1import torch2import torch.nn.functional as F3from torch import optim4from torch.nn import Module5from torchvision import models, transforms6from torchvision.datasets import ImageFolder7from PIL import Image8import numpy as np9import onnxruntime10import gradio as gr11import json12 13 14def get_image(x):15    return x.split(', ')[0]16 17 18def to_numpy(tensor):19    return tensor.detach().cpu().numpy() if tensor.requires_grad else tensor.cpu().numpy()20 21 22# Transform image to ToTensor23def transform_image(myarray):24    transform = transforms.Compose([25        transforms.Resize(224),26        transforms.CenterCrop(224),27        transforms.ToTensor(),28        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),29        ])30    image = Image.fromarray(np.uint8(myarray)).convert('RGB')31    image = transform(image).unsqueeze(0)32    return image33 34 35f = open('imagenet_label.json',)36label_map=json.load(f)37f.close()38 39# Load list of images for similarity40sub_test_list = open('img_list.txt', 'r')41sub_test_list = [i.strip() for i in sub_test_list]42 43# Load images embedding for similarity44embeddings = torch.load('embeddings.pt')45 46# Configure47options = onnxruntime.SessionOptions()48options.intra_op_num_threads = 849options.inter_op_num_threads = 850 51# Load model52PATH = 'model_onnx.onnx'53ort_session = onnxruntime.InferenceSession(PATH, sess_options=options)54input_name = ort_session.get_inputs()[0].name55 56 57# predict multi-level classification58def get_classification(img):59 60    image_tensor = transform_image(img)61    ort_inputs = {input_name: to_numpy(image_tensor)}62    x = ort_session.run(None, ort_inputs)63    predictions = torch.topk(torch.from_numpy(x[0]), k=5).indices.squeeze(0).tolist()64 65    result = {}66    for i in predictions:67        label = label_map[str(i)]68        prob = x[0][0, i].item()69        result[label] = prob70    return result71 72 73iface = gr.Interface(74    get_classification,75    gr.inputs.Image(shape=(200, 200)),76    outputs="label",77    title = 'Universal Image Classification',78    description = "Imagenet classification from Mobilenetv3 converting to ONNX runtime",79    article = "Author: <a href=\"https://huggingface.co/vumichien\">Vu Minh Chien</a>.",80)81iface.launch()82