SV12/ERA_Session13
0
1import gradio as gr2import random3import numpy as np4from PIL import Image5import torch6import torchvision7 8from pytorch_grad_cam import GradCAM9from pytorch_grad_cam.utils.image import show_cam_on_image10 11from models.resnet_lightning import ResNet12from utils.data import CIFARDataModule13from utils.transforms import test_transform14from utils.common import get_misclassified_data15 16inv_normalize = torchvision.transforms.Normalize(17 mean=[-0.50 / 0.23, -0.50 / 0.23, -0.50 / 0.23], std=[1 / 0.23, 1 / 0.23, 1 / 0.23]18)19 20datamodule = CIFARDataModule()21datamodule.setup()22classes = datamodule.train_dataset.classes23 24model = ResNet.load_from_checkpoint("model.ckpt")25model = model.to("cpu")26 27prediction_image = None28 29 30def upload_file(files):31 file_paths = [file.name for file in files]32 return file_paths33 34 35def read_image(path):36 img = Image.open(path)37 img.load()38 data = np.asarray(img, dtype="uint8")39 return data40 41 42def sample_images():43 images = []44 length = len(datamodule.test_dataset)45 classes = datamodule.train_dataset.classes46 for i in range(10):47 idx = random.randint(0, length - 1)48 image, label = datamodule.test_dataset[idx]49 image = inv_normalize(image).permute(1, 2, 0).numpy()50 images.append((image, classes[label]))51 return images52 53 54def get_misclassified_images(misclassified_count):55 misclassified_images = []56 misclassified_data = get_misclassified_data(57 model=model,58 device="cpu",59 test_loader=datamodule.test_dataloader(),60 count=misclassified_count,61 )62 for i in range(misclassified_count):63 img = misclassified_data[i][0].squeeze().to("cpu")64 img = inv_normalize(img)65 img = np.transpose(img.numpy(), (1, 2, 0))66 label = f"Label: {classes[misclassified_data[i][1].item()]} | Prediction: {classes[misclassified_data[i][2].item()]}"67 misclassified_images.append((img, label))68 return misclassified_images69 70 71def get_gradcam_images(gradcam_layer, gradcam_count, gradcam_opacity):72 gradcam_images = []73 if gradcam_layer == "Layer1":74 target_layers = [model.layer1[-1]]75 elif gradcam_layer == "Layer2":76 target_layers = [model.layer2[-1]]77 else:78 target_layers = [model.layer3[-1]]79 80 cam = GradCAM(model=model, target_layers=target_layers, use_cuda=False)81 data = get_misclassified_data(82 model=model,83 device="cpu",84 test_loader=datamodule.test_dataloader(),85 count=gradcam_count,86 )87 for i in range(gradcam_count):88 input_tensor = data[i][0]89 90 # Get the activations of the layer for the images91 grayscale_cam = cam(input_tensor=input_tensor, targets=None)92 grayscale_cam = grayscale_cam[0, :]93 94 # Get back the original image95 img = input_tensor.squeeze(0).to("cpu")96 if inv_normalize is not None:97 img = inv_normalize(img)98 rgb_img = np.transpose(img, (1, 2, 0))99 rgb_img = rgb_img.numpy()100 101 # Mix the activations on the original image102 visualization = show_cam_on_image(103 rgb_img, grayscale_cam, use_rgb=True, image_weight=gradcam_opacity104 )105 label = f"Label: {classes[data[i][1].item()]} | Prediction: {classes[data[i][2].item()]}"106 gradcam_images.append((visualization, label))107 return gradcam_images108 109 110def show_hide_misclassified(status):111 if not status:112 return {misclassified_count: gr.update(visible=False)}113 return {misclassified_count: gr.update(visible=True)}114 115 116def show_hide_gradcam(status):117 if not status:118 return [gr.update(visible=False) for i in range(3)]119 return [gr.update(visible=True) for i in range(3)]120 121 122def set_prediction_image(evt: gr.SelectData, gallery):123 global prediction_image124 if isinstance(gallery[evt.index], dict):125 prediction_image = gallery[evt.index]["name"]126 else:127 prediction_image = gallery[evt.index][0]["name"]128 129 130def predict(131 is_misclassified,132 misclassified_count,133 is_gradcam,134 gradcam_count,135 gradcam_layer,136 gradcam_opacity,137 num_classes,138):139 misclassified_images = None140 if is_misclassified:141 misclassified_images = get_misclassified_images(int(misclassified_count))142 143 gradcam_images = None144 if is_gradcam:145 gradcam_images = get_gradcam_images(146 gradcam_layer, int(gradcam_count), gradcam_opacity147 )148 149 img = read_image(prediction_image)150 image_transformed = test_transform(image=img)["image"]151 output = model(image_transformed.unsqueeze(0))152 preds = torch.softmax(output, dim=1).squeeze().detach().numpy()153 indices = (154 output.argsort(descending=True).squeeze().detach().numpy()[: int(num_classes)]155 )156 predictions = {classes[i]: round(float(preds[i]), 2) for i in indices}157 158 return {159 miscalssfied_output: gr.update(value=misclassified_images),160 gradcam_output: gr.update(value=gradcam_images),161 prediction_label: gr.update(value=predictions),162 }163 164 165with gr.Blocks() as app:166 gr.Markdown("## ERA Session13 - CIFAR10 Classification with ResNet")167 with gr.Row():168 with gr.Column():169 with gr.Group():170 is_misclassified = gr.Checkbox(171 label="Misclassified Images", info="Display misclassified images?"172 )173 misclassified_count = gr.Dropdown(174 choices=["10", "20"],175 label="Select Number of Images",176 info="Number of Misclassified images",177 visible=False,178 interactive=True,179 )180 is_misclassified.input(181 show_hide_misclassified,182 inputs=[is_misclassified],183 outputs=[misclassified_count],184 )185 with gr.Group():186 is_gradcam = gr.Checkbox(187 label="GradCAM Images",188 info="Display GradCAM images?",189 )190 gradcam_count = gr.Dropdown(191 choices=["10", "20"],192 label="Select Number of Images",193 info="Number of GradCAM images",194 interactive=True,195 visible=False,196 )197 gradcam_layer = gr.Dropdown(198 choices=["Layer1", "Layer2", "Layer3"],199 label="Select the layer",200 info="Please select the layer for which the GradCAM is required",201 interactive=True,202 visible=False,203 )204 gradcam_opacity = gr.Slider(205 minimum=0,206 maximum=1,207 value=0.6,208 label="Opacity",209 info="Opacity of GradCAM output",210 interactive=True,211 visible=False,212 )213 214 is_gradcam.input(215 show_hide_gradcam,216 inputs=[is_gradcam],217 outputs=[gradcam_count, gradcam_layer, gradcam_opacity],218 )219 with gr.Group():220 # file_output = gr.File(file_types=["image"])221 with gr.Group():222 upload_gallery = gr.Gallery(223 value=None,224 label="Uploaded images",225 show_label=False,226 elem_id="gallery_upload",227 columns=5,228 rows=2,229 height="auto",230 object_fit="contain",231 )232 upload_button = gr.UploadButton(233 "Click to Upload images",234 file_types=["image"],235 file_count="multiple",236 )237 upload_button.upload(upload_file, upload_button, upload_gallery)238 239 with gr.Group():240 sample_gallery = gr.Gallery(241 value=sample_images,242 label="Sample images",243 show_label=True,244 elem_id="gallery_sample",245 columns=5,246 rows=2,247 height="auto",248 object_fit="contain",249 )250 251 upload_gallery.select(set_prediction_image, inputs=[upload_gallery])252 sample_gallery.select(set_prediction_image, inputs=[sample_gallery])253 254 with gr.Group():255 num_classes = gr.Dropdown(256 choices=[str(i + 1) for i in range(10)],257 label="Select Number of Top Classes",258 info="Number of Top target classes to be shown",259 )260 run_btn = gr.Button()261 with gr.Column():262 with gr.Group():263 miscalssfied_output = gr.Gallery(264 value=None, label="Misclassified Images", show_label=True265 )266 with gr.Group():267 gradcam_output = gr.Gallery(268 value=None, label="GradCAM Images", show_label=True269 )270 with gr.Group():271 prediction_label = gr.Label(value=None, label="Predictions")272 273 run_btn.click(274 predict,275 inputs=[276 is_misclassified,277 misclassified_count,278 is_gradcam,279 gradcam_count,280 gradcam_layer,281 gradcam_opacity,282 num_classes,283 ],284 outputs=[miscalssfied_output, gradcam_output, prediction_label],285 )286 287 288app.launch(server_name="0.0.0.0", server_port=8000)289 