Samarth991/GradCam_Explainable-AI
0
1import torch.nn as nn2import torch.nn.functional as F3import gradio as gr4import torch5import random6from collections import OrderedDict7from pytorch_grad_cam import GradCAM8from pytorch_grad_cam.utils.image import show_cam_on_image9from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget10import numpy as np11from PIL import Image12from torchvision import transforms13 14dropout_value = 0.115class ResBlock(nn.Module):16 def __init__(self, in_channels, out_channels):17 super(ResBlock,self).__init__()18 self.res_block = nn.Sequential(19 nn.Conv2d(in_channels=in_channels, out_channels = out_channels, kernel_size=3, stride =1 , padding =1),20 nn.BatchNorm2d(out_channels),21 nn.ReLU(),22 nn.Conv2d(in_channels=out_channels, out_channels = out_channels, kernel_size=3, stride =1 , padding =1),23 nn.BatchNorm2d(out_channels),24 nn.ReLU(),25 )26 27 def forward (self, x):28 x = self.res_block(x)29 return x30 31 32class LayerBlock(nn.Module):33 def __init__(self, in_channels, out_channels):34 super(LayerBlock,self).__init__()35 self.layer_block = nn.Sequential(36 nn.Conv2d(in_channels=in_channels, out_channels = out_channels, kernel_size=3, stride =1 , padding =1),37 nn.MaxPool2d(kernel_size=2,stride=2),38 nn.BatchNorm2d(out_channels),39 nn.ReLU(),40 )41 42 def forward (self, x):43 x = self.layer_block(x)44 return x45 46class custom_resnet_s10(nn.Module):47 def __init__(self, num_classes=10):48 super(custom_resnet_s10,self).__init__()49 50 self.PrepLayer = nn.Sequential(51 nn.Conv2d(in_channels = 3, out_channels=64, kernel_size = 3, stride = 1, padding =1),52 nn.BatchNorm2d(64),53 nn.ReLU(),54 )55 self.Layer1 = LayerBlock(in_channels = 64, out_channels=128)56 self.resblock1 = ResBlock(in_channels =128, out_channels=128)57 self.Layer2 = LayerBlock(in_channels = 128, out_channels=256)58 self.resblock2 = ResBlock(in_channels =256, out_channels=256)59 self.Layer3 = LayerBlock(in_channels = 256, out_channels=512)60 self.resblock3 = ResBlock(in_channels =512, out_channels=512)61 self.max_pool4 = nn.MaxPool2d(kernel_size=4, stride=4) # 512,512, 4/4 = 512,512,162 self.fc = nn.Linear(512,num_classes)63 64 def forward(self,x):65 x = self.PrepLayer(x)66 x = self.Layer1(x)67 resl1 = self.resblock1(x)68 69 x = x+resl170 x = self.Layer2(x)71 resl2 = self.resblock2(x)72 x = x+resl273 74 75 x = self.Layer3(x)76 77 resl3 = self.resblock3(x)78 x = x+resl379 x = self.max_pool4(x)80 x = x.view(x.size(0),-1)81 82 x = self.fc(x)83 return x84 85def get_device():86 if torch.cuda.is_available():87 device = "cuda"88 elif torch.backends.mps.is_available():89 device = "mps"90 else:91 device = "cpu"92 print("Device Selected:", device)93 return device94 95DEVICE = get_device()96 97# Load the list of tensors from the file98loaded_misclassified_image_list = torch.load('misclassified_images_list.pt')99 100# Instantiate the model (make sure it has the same architecture)101loaded_model = custom_resnet_s10()102loaded_model = loaded_model.to(DEVICE)103 104# Load the saved state dictionary105loaded_model.load_state_dict(torch.load('model.pth', map_location=DEVICE), strict=False)106 107# Put the loaded model in evaluation mode108loaded_model.eval()109 110classes = ['plane', 'car', 'bird', 'cat', 'deer','dog', 'frog', 'horse', 'ship', 'truck']111mean = (0.49139968, 0.48215827, 0.44653124)112std = (0.24703233, 0.24348505, 0.26158768)113transform = transforms.Compose([114 transforms.ToTensor(),115 transforms.Normalize(mean=mean, std=std)116])117 118dict_layer = {'layer3': loaded_model.resblock2.res_block[-1],119 'layer4': loaded_model.resblock3.res_block[-1]}120 121def view_gradcam_images(choice_gradcam):122 if choice_gradcam == "Yes (View Existing Images)":123 return gr.update(label ="Number of GradCAM Images to view", visible=True, interactive = True), \124 gr.update(visible=True), \125 gr.update(visible=True), gr.update(visible=True), \126 gr.update(visible=False) # Gallery not shown as yet127 else:128 #TODO: to be completed129 return gr.update(visible=False), gr.update(visible=False),gr.update(visible=False),gr.update(visible=False),gr.update(visible=False)130 131def process_gradcam_images(num_images,layer,opacity,image_list=None):132 133 if not image_list:134 selected_data = random.sample(loaded_misclassified_image_list, min(num_images,len(loaded_misclassified_image_list)))135 else:136 selected_data = [image_list]137 138 layer_model = dict_layer.get(layer)139 cam = GradCAM(model=loaded_model, target_layers = [layer_model], use_cuda = False)140 grad_images = []141 inv_normalize = transforms.Normalize(142 mean=[-0.50/0.2197, -0.50/0.1858, -0.50/0.1569], # mean_ds = [0.2197, 0.1858, 0.1569]143 std=[1/0.1810, 1/0.1635, 1/0.1511] # std_dev_ds =[0.1810, 0.1635, 0.1511]144 )145 for i, (img, pred, correct) in enumerate(selected_data):146 input_tensor = img.unsqueeze(0)147 targets = [ClassifierOutputTarget(pred)]148 grayscale_cam = cam(input_tensor=input_tensor, targets=targets)149 150 grayscale_cam = grayscale_cam[0, :]151 152 # Get back the original image153 img = input_tensor.squeeze(0).to('cpu')154 img = inv_normalize(img)155 rgb_img = np.transpose(img, (1, 2, 0))156 rgb_img = torch.clamp(rgb_img, max = 1)157 rgb_img = rgb_img.numpy()158 159 visualization = show_cam_on_image(rgb_img, grayscale_cam, use_rgb=True, image_weight=opacity)160 if not image_list:161 grad_images.append(((visualization),f'Pred: {classes[pred.cpu()]} | Truth :{classes[correct.cpu()]}'))162 else:163 grad_images.append(((visualization),f'Prediction: {classes[pred.cpu()]}'))164 165 print(str(num_images) + "**" + str(layer) + "**" + str(opacity))166 return grad_images, gr.update(visible=True)167 168 169 170def process_misclassified_images(num_images):171 selected_data = random.sample(loaded_misclassified_image_list, min(num_images,len(loaded_misclassified_image_list)))172 misclassified_images = []173 for i, (img, pred, correct) in enumerate(selected_data):174 img, pred, target = img.cpu().numpy().astype(dtype=np.float32), pred.cpu(), correct.cpu()175 for j in range(img.shape[0]):176 img[j] = (img[j] * std[j]) + mean[j]177 img = np.transpose(img, (1, 2, 0))178 img = Image.fromarray((img * 255).astype(np.uint8))179 misclassified_images.append(((img),f'Pred: {classes[pred]} | Truth :{classes[correct]}'))180 return misclassified_images, gr.update(visible=True)181 182def view_misclassified_images(choice_misclassified):183 if choice_misclassified == "Yes":184 return gr.update(label ="Number of Misclassified Images to view", visible=True, interactive = True),gr.update(visible=True),gr.update(visible=False)185 else:186 return gr.update(visible=False),gr.update(visible=False),gr.update(visible=False)187 188def classify_image(image, num_classes=3, grad_cam_choice = False, layer = None, opacity = 0.8 ):189 # transforming image and getting prediction from model190 transformed_image = transform(image)191 image_tensor = transformed_image.to(DEVICE).unsqueeze(0)#transform(torch.tensor(image).to(DEVICE)).unsqueeze(0) # making it a batch192 193 # sending it to model to get prediction194 logits = loaded_model(image_tensor) # logits195 output = F.softmax(logits.view(-1)) #F.softmax(output.flatten(), dim=-1) #196 197 confidences = [(classes[i], float(output[i])) for i in range(len(classes))]198 confidences.sort(key=lambda x: x[1], reverse=True)199 confidences = OrderedDict(confidences[:num_classes])200 label = torch.argmax(output).item()201 202 if grad_cam_choice:203 print("** Before Calling **",transformed_image.shape)204 image_list = [transformed_image.to(DEVICE),torch.tensor(label).to(DEVICE),torch.tensor(label).to(DEVICE)]205 206 grad_cam_output,_ = process_gradcam_images(num_images = 1,layer = layer,opacity= opacity,image_list=image_list)207 208 return confidences, grad_cam_output , gr.update(visible=True)209 else:210 return confidences, gr.update(visible=False),gr.update(visible=False)211 212 213with gr.Blocks() as demo:214 with gr.Tab("GradCam"):215 gr.Markdown(216 """217 Visualize Class Activations Maps (helps to see what the model is actually looking at in the image) generated by the model's layer for the predicted class218 - For existing images219 - For new images (choose an example image or upload your own)220 """221 )222 with gr.Column():223 with gr.Box():224 radio_gradcam = gr.Radio(["Yes (View Existing Images)", "No (New or Example Images)"], label="Do you want to view existing GradCAM images?")225 with gr.Column():226 with gr.Row():227 slider_gradcam_num_images = gr.Slider(minimum=1, maximum =10, value = 1, step =1, visible= False, interactive = False)228 dropdown_gradcam_layer = gr.Dropdown(choices=['layer4', 'layer3'], value = "layer4", label="Please select the layer from which the GradCAM would be taken", interactive = True, visible= False)229 slider_gradcam_opacity = gr.Slider(label ="Opacity of Images", minimum=0.05, maximum =1.00, value = 0.70, step =0.05, visible= False, interactive = True)230 button_gradcam = gr.Button("View GradCAM Output", visible = False)231 # txt_gradcam = gr.Textbox ("GradCAM output here" , visible = True)232 output_gallery_gradcam=gr.Gallery(label="GradCAM Output", min_width=512,columns=4, visible = False)233 with gr.Box():234 with gr.Row():235 with gr.Column():236 input_image_classify = gr.Image(label="Classification",type="pil", shape=(32, 32))237 slider_classify_num_classes = gr.Slider(label="Select the number of top classes to be shown",minimum=1, maximum =10, value = 3, step = 1, visible= True, interactive = True)238 checkbox_gradcam_classify = gr.Checkbox(label="Enable GradCAM", value=True, info="Do you want to see Class Activation Maps?", visible=True)239 # txt_classify= gr.Textbox ("Classification output here" , visible = True)240 dropdown_gradcam_classify_layer = gr.Dropdown(choices=['layer4', 'layer3'], value = "layer4", label="Please select the layer from which the GradCAM would be taken", interactive = True, visible= True)241 slider_gradcam_classify_opacity = gr.Slider(label ="Opacity of Images", minimum=0.05, maximum =1.00, value = 0.80, step =0.05, visible= True, interactive = True)242 button_classify = gr.Button("Submit to Classify Image", visible = True)243 244 with gr.Column():245 label_classify = gr.Label(num_top_classes=10, visible = True)246 gallery_gradcam_classify = gr.Gallery(label="GradCAM Output", min_width=256,columns=1, visible = True)247 with gr.Row():248 gr.Examples(['bird1.jpg','car1.jpg','deer1.jpg','frog1.jpg','plane1.jpg','ship1.jpg','truck1.jpg',"cat1.jpg","dog1.jpg","horse1.jpg"],inputs=[input_image_classify])249 with gr.Tab("Misclassified Examples"):250 gr.Markdown(251 """252 The AI model is not able to predict correct image labels all the time.253 254 Select "Yes" to visualize the misclassified images with their model predicted label and ground truth label.255 """256 )257 with gr.Column():258 with gr.Box():259 radio_misclassified = gr.Radio(["Yes", "No"], label="Do you want to view Misclassified images?")260 slider_misclassified_num_images = gr.Slider(minimum=1, maximum =10, value = 1, step =1, visible= False, interactive = False)261 button_misclassified = gr.Button("View Misclassified Output", visible = False)262 # txt_misclassified = gr.Textbox ("Misclassified output here" , visible = True)263 output_gallery_misclassification=gr.Gallery(label="Misclassification Output (Predicted/Truth)", min_width=512,columns=5, visible = False)264 265 radio_gradcam.change(fn=view_gradcam_images, inputs=radio_gradcam, outputs=[slider_gradcam_num_images, dropdown_gradcam_layer,slider_gradcam_opacity,button_gradcam, output_gallery_gradcam])266 button_gradcam.click(fn = process_gradcam_images, inputs = [slider_gradcam_num_images,dropdown_gradcam_layer,slider_gradcam_opacity], outputs = [output_gallery_gradcam,output_gallery_gradcam])267 268 radio_misclassified.change(fn=view_misclassified_images, inputs=radio_misclassified, outputs=[slider_misclassified_num_images,button_misclassified,output_gallery_misclassification])269 button_misclassified.click(fn = process_misclassified_images, inputs = [slider_misclassified_num_images], outputs = [output_gallery_misclassification,output_gallery_misclassification])270 button_classify.click(fn=classify_image, inputs =[input_image_classify,slider_classify_num_classes,checkbox_gradcam_classify,dropdown_gradcam_classify_layer,slider_gradcam_classify_opacity], outputs = [label_classify,gallery_gradcam_classify,gallery_gradcam_classify])271demo.launch ()272 273 