all-things-vits/comparative-explainability
0
1import sys2 3import cv24import numpy as np5import torch6 7from imagenet_class_indices import CLS2IDX8 9sys.path.append("Transformer-Explainability")10 11 12from baselines.ViT.ViT_explanation_generator import LRP, Baselines13from baselines.ViT.ViT_LRP import vit_base_patch16_224 as vit_LRP14from baselines.ViT.ViT_new import vit_base_patch16_224 as vit15 16 17# create heatmap from mask on image18def show_cam_on_image(img, mask):19 heatmap = cv2.applyColorMap(np.uint8(255 * mask), cv2.COLORMAP_JET)20 heatmap = np.float32(heatmap) / 25521 cam = heatmap + np.float32(img)22 cam = cam / np.max(cam)23 return cam24 25 26# initialize ViT pretrained27model = vit_LRP(pretrained=True)28model.eval()29attribution_generator = LRP(model)30model_baseline = vit(pretrained=True)31model_baseline.eval()32baselines_generator = Baselines(model_baseline)33 34 35def generate_visualization(36 original_image, class_index=None, method="transformer_attribution", LRP=True37):38 if LRP:39 transformer_attribution = attribution_generator.generate_LRP(40 original_image.unsqueeze(0), method=method, index=class_index41 ).detach()42 else:43 if method == "gradcam":44 transformer_attribution = baselines_generator.generate_cam_attn(45 original_image.unsqueeze(0), index=class_index46 ).detach()47 else:48 transformer_attribution = baselines_generator.generate_rollout(49 original_image.unsqueeze(0)50 ).detach()51 if method != "full":52 transformer_attribution = transformer_attribution.reshape(1, 1, 14, 14)53 transformer_attribution = torch.nn.functional.interpolate(54 transformer_attribution, scale_factor=16, mode="bilinear"55 )56 else:57 transformer_attribution = transformer_attribution.reshape(1, 1, 224, 224)58 transformer_attribution = (59 transformer_attribution.reshape(224, 224).data.cpu().numpy()60 )61 transformer_attribution = (62 transformer_attribution - transformer_attribution.min()63 ) / (transformer_attribution.max() - transformer_attribution.min())64 65 image_transformer_attribution = original_image.permute(1, 2, 0).data.cpu().numpy()66 image_transformer_attribution = (67 image_transformer_attribution - image_transformer_attribution.min()68 ) / (image_transformer_attribution.max() - image_transformer_attribution.min())69 vis = show_cam_on_image(image_transformer_attribution, transformer_attribution)70 vis = np.uint8(255 * vis)71 vis = cv2.cvtColor(np.array(vis), cv2.COLOR_RGB2BGR)72 return vis73 74 75def print_top_classes(predictions, **kwargs):76 # Print Top-5 predictions77 prob = torch.softmax(predictions, dim=1)78 class_indices = predictions.data.topk(5, dim=1)[1][0].tolist()79 max_str_len = 080 class_names = []81 for cls_idx in class_indices:82 class_names.append(CLS2IDX[cls_idx])83 if len(CLS2IDX[cls_idx]) > max_str_len:84 max_str_len = len(CLS2IDX[cls_idx])85 86 print("Top 5 classes:")87 for cls_idx in class_indices:88 output_string = "\t{} : {}".format(cls_idx, CLS2IDX[cls_idx])89 output_string += " " * (max_str_len - len(CLS2IDX[cls_idx])) + "\t\t"90 output_string += "value = {:.3f}\t prob = {:.1f}%".format(91 predictions[0, cls_idx], 100 * prob[0, cls_idx]92 )93 print(output_string)94 