HugoHE/gtsrb-classifier
0
1import torch2from torch.nn import functional as F3import torchvision4from torchvision import transforms, models5import pytorch_lightning as pl6from pytorch_lightning import LightningModule, Trainer7from PIL import Image8import gradio as gr9 10classes = ['Speed limit (20km/h)',11 'Speed limit (30km/h)',12 'Speed limit (50km/h)',13 'Speed limit (60km/h)',14 'Speed limit (70km/h)',15 'Speed limit (80km/h)',16 'End of speed limit (80km/h)',17 'Speed limit (100km/h)',18 'Speed limit (120km/h)',19 'No passing',20 'No passing veh over 3.5 tons',21 'Right-of-way at intersection',22 'Priority road',23 'Yield',24 'Stop',25 'No vehicles',26 'Veh > 3.5 tons prohibited',27 'No entry',28 'General caution',29 'Dangerous curve left',30 'Dangerous curve right',31 'Double curve',32 'Bumpy road',33 'Slippery road',34 'Road narrows on the right',35 'Road work',36 'Traffic signals',37 'Pedestrians',38 'Children crossing',39 'Bicycles crossing',40 'Beware of ice/snow',41 'Wild animals crossing',42 'End speed + passing limits',43 'Turn right ahead',44 'Turn left ahead',45 'Ahead only',46 'Go straight or right',47 'Go straight or left',48 'Keep right',49 'Keep left',50 'Roundabout mandatory',51 'End of no passing',52 'End no passing veh > 3.5 tons']53 54class LitGTSRB(pl.LightningModule):55 def __init__(self):56 super().__init__()57 self.model = models.resnet18(pretrained=False, num_classes=43)58 59 def forward(self, x):60 out = self.model(x)61 return F.log_softmax(out, dim=1)62 63def predict_image(image):64 model = LitGTSRB().load_from_checkpoint('resnet18.ckpt')65 model.eval()66 image = image.convert('RGB')67 test_transforms = transforms.Compose([68 transforms.Resize([224, 224]),69 transforms.ToTensor(), 70 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])71 ])72 image_tensor = test_transforms(image).float()73 image_tensor = image_tensor.unsqueeze_(0)74 with torch.no_grad():75 output = model(image_tensor)76 probs = torch.exp(output.data.cpu().squeeze())77 prediction_score , pred_label_idx = torch.topk(probs,5)78 class_top5 = [classes[idx] for idx in pred_label_idx.numpy()]79 return dict(zip(class_top5, map(float, prediction_score.numpy())))80image = gr.Image(type='pil')81label = gr.Label()82examples = ['1.png', '2.png', '3.png', '4.png', '5.png', '6.png', '7.png', '8.png']83intf = gr.Interface(fn=predict_image, inputs=image, outputs=label, examples=examples)84intf.launch(inline=True)