CoolFace
Apppublic

nandeesh-n/CIFAR10_Custom_ResNet

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
app.py137 linesDownload Raw Back to root
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