pytorch/EfficientNet
1
1import torch2import gradio as gr3import torchvision.transforms as transforms4 5device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")6 7efficientnet = torch.hub.load('NVIDIA/DeepLearningExamples:torchhub', 'nvidia_efficientnet_b0', pretrained=True)8utils = torch.hub.load('NVIDIA/DeepLearningExamples:torchhub', 'nvidia_convnets_processing_utils')9 10efficientnet.eval().to(device)11 12def inference(img):13 14 img_transforms = transforms.Compose(15 [transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor()]16 )17 18 img = img_transforms(img)19 with torch.no_grad():20 # mean and std are not multiplied by 255 as they are in training script21 # torch dataloader reads data into bytes whereas loading directly22 # through PIL creates a tensor with floats in [0,1] range23 mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)24 std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)25 img = img.float()26 img = img.unsqueeze(0).sub_(mean).div_(std)27 28 batch = torch.cat(29 [img]30 ).to(device)31 with torch.no_grad():32 output = torch.nn.functional.softmax(efficientnet(batch), dim=1)33 34 35 results = utils.pick_n_best(predictions=output, n=5)36 37 return results38 39title="EfficientNet"40description="Gradio demo for EfficientNet,EfficientNets are a family of image classification models, which achieve state-of-the-art accuracy, being an order-of-magnitude smaller and faster. Trained with mixed precision using Tensor Cores. To use it, simply upload your image or click on one of the examples below. Read more at the links below"41 42article = "<p style='text-align: center'><a href='https://arxiv.org/abs/1905.11946'>EfficientNet: Rethinking Model Scaling for Convolutional Neural Networks</a> | <a href='https://github.com/NVIDIA/DeepLearningExamples/tree/master/PyTorch/Classification/ConvNets/efficientnet'>Github Repo</a></p>"43 44examples=[['food.jpeg']]45gr.Interface(inference,gr.inputs.Image(type="pil"),"text",title=title,description=description,article=article,examples=examples).launch()