pytorch/AlexNet
3
1import os2import torch3import gradio as gr4from PIL import Image5from torchvision import transforms6 7torch.hub.download_url_to_file("https://github.com/pytorch/hub/raw/master/images/dog.jpg", "dog.jpg")8 9model = torch.hub.load('pytorch/vision:v0.9.0', 'alexnet', pretrained=True)10model.eval()11 12# Download ImageNet labels13os.system("wget https://raw.githubusercontent.com/pytorch/hub/master/imagenet_classes.txt")14 15def inference(input_image):16 17 preprocess = transforms.Compose([18 transforms.Resize(256),19 transforms.CenterCrop(224),20 transforms.ToTensor(),21 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),22 ])23 input_tensor = preprocess(input_image)24 input_batch = input_tensor.unsqueeze(0) # create a mini-batch as expected by the model25 26 # move the input and model to GPU for speed if available27 if torch.cuda.is_available():28 input_batch = input_batch.to('cuda')29 model.to('cuda')30 31 with torch.no_grad():32 output = model(input_batch)33 # The output has unnormalized scores. To get probabilities, you can run a softmax on it.34 probabilities = torch.nn.functional.softmax(output[0], dim=0)35 # Read the categories36 with open("imagenet_classes.txt", "r") as f:37 categories = [s.strip() for s in f.readlines()]38 # Show top categories per image39 top5_prob, top5_catid = torch.topk(probabilities, 5)40 result = {}41 for i in range(top5_prob.size(0)):42 result[categories[top5_catid[i]]] = top5_prob[i].item()43 return result44 45inputs = gr.inputs.Image(type='pil')46outputs = gr.outputs.Label(type="confidences",num_top_classes=5)47 48title = "ALEXNET"49description = "Gradio demo for Alexnet, the 2012 ImageNet winner achieved a top-5 error of 15.3%, more than 10.8 percentage points lower than that of the runner up. To use it, simply upload your image, or click one of the examples to load them. Read more at the links below."50article = "<p style='text-align: center'><a href='https://arxiv.org/abs/1404.5997'>One weird trick for parallelizing convolutional neural networks</a> | <a href='https://github.com/pytorch/vision/blob/master/torchvision/models/alexnet.py'>Github Repo</a></p>"51 52examples = [53 ['dog.jpg']54]55gr.Interface(inference, inputs, outputs, title=title, description=description, article=article, examples=examples, analytics_enabled=False).launch()