nandeesh-n/CIFAR10_Custom_ResNet
0
1# Import libraries2import torch3from torchvision import transforms4import numpy as np5from pytorch_grad_cam import GradCAM6from pytorch_grad_cam.utils.image import show_cam_on_image7from models.custom_resnet import LitCustomResNet8import gradio as gr9import cv210import os11 12examples_dir = os.path.join(os.path.dirname(__file__), 'examples')13examples = [[os.path.join(examples_dir, img), img.split('.')[0].split('_')[0]] for img in os.listdir(examples_dir)]14classes = ['plane', 'car', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck']15model = LitCustomResNet(in_ch=3)16model.load_state_dict(torch.load('model.pth', map_location=torch.device('cpu')), strict=False)17 18inv_normalize = transforms.Normalize(19 mean=[-0.50/0.23, -0.50/0.23, -0.50/0.23],20 std=[1/0.23, 1/0.23, 1/0.23]21)22 23def inference(input_img, opacity, top_n, disp_grad, num_grad, gc_layer, disp_mis, num_mis):24 transform = transforms.ToTensor()25 _ = np.random.shuffle(examples)26 gc_vis_imgs = []27 mis_vis_imgs = []28 29 _, confidences, gc_vis_img = get_pred_vis(input_img, transform, gc_layer, opacity)30 confidences = dict(sorted(confidences.items(), key=lambda item: item[1], reverse=True)[:top_n])31 gc_vis_imgs.append(gc_vis_img)32 33 # Display GradCAM images from the examples34 if disp_grad:35 for i in range(num_grad):36 ex_img = cv2.imread(examples[i][0])37 ex_img = cv2.cvtColor(ex_img, cv2.COLOR_BGR2RGB)38 _, _, ex_gc_vis_img = get_pred_vis(ex_img, transform, gc_layer, opacity)39 gc_vis_imgs.append([ex_gc_vis_img, f'Actual: {examples[i][1]}'])40 41 # Display misclassified images from the examples42 if disp_mis:43 for i in range(num_mis):44 ex_img = cv2.imread(examples[i][0])45 ex_img = cv2.cvtColor(ex_img, cv2.COLOR_BGR2RGB)46 pred, _, ex_gc_vis_img = get_pred_vis(ex_img, transform, gc_layer, opacity)47 mis_vis_imgs.append([ex_img, f'Actual: {examples[i][1]}; Pred: {pred}'])48 return confidences, gc_vis_imgs, mis_vis_imgs49 50def get_pred_vis(input_img, transform, gc_layer, opacity):51 org_img = input_img52 input_img = transform(input_img)53 input_img = input_img54 input_img = input_img.unsqueeze(0)55 outputs = model(input_img)56 softmax = torch.nn.Softmax(dim=0)57 o = softmax(outputs.flatten())58 confidences = {classes[i]: float(o[i]) for i in range(10)}59 target_layers = [model.layer2[gc_layer]]60 cam = GradCAM(model=model, target_layers=target_layers, use_cuda=False)61 grayscale_cam = cam(input_tensor=input_img, targets=None)62 grayscale_cam = grayscale_cam[0, :]63 img = input_img.squeeze(0)64 img = inv_normalize(img)65 rgb_img = np.transpose(img, (1, 2, 0))66 rgb_img = rgb_img.numpy()67 visualization = show_cam_on_image(org_img/255, grayscale_cam, use_rgb=True, image_weight=opacity)68 return classes[o.argmax().item()], confidences, visualization69 70# Creating the gradio application using `blocks`.71with gr.Blocks() as classification_app:72 gr.Markdown('<h1 style="text-align: center;">Image classification by custom ResNet on CIFAR10 dataset + GradCAM</h1>')73 with gr.Row():74 # Define inputs75 with gr.Column():76 with gr.Row():77 input_img = gr.Image(label='Input Image', shape=(32, 32))78 with gr.Row():79 top_n = gr.Slider(label='# of top classes', minimum=1, maximum=10, value=3, step=1)80 gr.Markdown('### GradCam Settings:')81 with gr.Row():82 disp_grad = gr.Radio(label='Display GradCam images?', choices=['Yes', 'No'], value='Yes')83 num_grad = gr.Slider(label='# of GradCam images to visualize?', maximum=10, value=1, step=1)84 with gr.Row():85 gc_layer = gr.Radio(label='Which layer?', choices=[-1, -2], value=-2)86 opacity = gr.Slider(label='Opacity', maximum=1, value=0.5, step=0.1)87 gr.Markdown('### Misclassification Settings:')88 with gr.Row():89 disp_mis = gr.Radio(label='Display misclassified images?', choices=['Yes', 'No'], value='No')90 num_mis = gr.Slider(label='# of misclassified images to display?', maximum=10, value=1, step=1, interactive=False)91 with gr.Row():92 btn = gr.Button('Predict')93 94 # Define outputs95 with gr.Column():96 out_label = gr.Label(label='Output Label')97 gc_vis = gr.Gallery(label='GradCAM', value=[], columns=3, object_fit='contain')98 mis_vis = gr.Gallery(label='Misclassified images', value=[], columns=3, object_fit='contain', visible=False)99 100 # Event listeners101 # Enable/ disable the components based on whether user wants to visualize GradCAM images.102 disp_grad.change(103 fn=lambda value: [gr.update(interactive=(value == 'Yes')),104 gr.update(interactive=(value == 'Yes')),105 gr.update(interactive=(value == 'Yes')),106 gr.update(visible=(value == 'Yes'))],107 inputs=disp_grad,108 outputs=[num_grad, gc_layer, opacity, gc_vis]109 )110 111 112 # Enable/ disable the components based on whether user wants to visualize misclassified images.113 disp_mis.change(114 fn=lambda value: [gr.update(interactive=(value == 'Yes')),115 gr.update(visible=(value == 'Yes'))],116 inputs=disp_mis,117 outputs=[num_mis, mis_vis]118 )119 120 # Button click handler121 btn.click(fn=inference, inputs=[122 input_img, opacity, top_n, disp_grad, num_grad, gc_layer, disp_mis, num_mis123 ],124 outputs=[125 out_label, gc_vis, mis_vis126 ])127 128 with gr.Row():129 true_lbl = gr.Textbox(label='True Label', visible=False)130 ex = gr.Examples(131 examples=examples,132 inputs=[input_img, true_lbl],133 outputs=[]134 )135 136classification_app.launch()137 