Team Ai
Apppublic

all-things-vits/comparative-explainability

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
generic_utils.py94 linesDownload Raw Back to root
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