CoolFace
Apppublic

amsterdamNLP/CLIP-attention-rollout

sourceHugging Faceafl-3.0updated 2y agoView on Hugging Face
3likes
1import torch2import CLIP.clip as clip3from PIL import Image4import numpy as np5import cv26import matplotlib.pyplot as plt7from captum.attr import visualization8import os9 10 11from CLIP.clip.simple_tokenizer import SimpleTokenizer as _Tokenizer12_tokenizer = _Tokenizer()13 14#@title Control context expansion (number of attention layers to consider)15#@title Number of layers for image Transformer16#start_layer =  11#@param {type:"number"}17 18#@title Number of layers for text Transformer19start_layer_text =  11#@param {type:"number"}20 21 22def interpret(image, texts, model, device, start_layer):23    batch_size = texts.shape[0]24    images = image.repeat(batch_size, 1, 1, 1)25    logits_per_image, logits_per_text = model(images, texts)26    probs = logits_per_image.softmax(dim=-1).detach().cpu().numpy()27    index = [i for i in range(batch_size)]28    one_hot = np.zeros((logits_per_image.shape[0], logits_per_image.shape[1]), dtype=np.float32)29    one_hot[torch.arange(logits_per_image.shape[0]), index] = 130    one_hot = torch.from_numpy(one_hot).requires_grad_(True)31    one_hot = torch.sum(one_hot.to(device) * logits_per_image)32    model.zero_grad()33 34    image_attn_blocks = list(dict(model.visual.transformer.resblocks.named_children()).values())35    num_tokens = image_attn_blocks[0].attn_probs.shape[-1]36    R = torch.eye(num_tokens, num_tokens, dtype=image_attn_blocks[0].attn_probs.dtype).to(device)37    R = R.unsqueeze(0).expand(batch_size, num_tokens, num_tokens)38    for i, blk in enumerate(image_attn_blocks):39        if i < start_layer:40            continue41        grad = torch.autograd.grad(one_hot, [blk.attn_probs], retain_graph=True)[0].detach()42        cam = blk.attn_probs.detach()43        cam = cam.reshape(-1, cam.shape[-1], cam.shape[-1])44        grad = grad.reshape(-1, grad.shape[-1], grad.shape[-1])45        cam = grad * cam46        cam = cam.reshape(batch_size, -1, cam.shape[-1], cam.shape[-1])47        cam = cam.clamp(min=0).mean(dim=1)48        R = R + torch.bmm(cam, R)49    image_relevance = R[:, 0, 1:]50 51 52    text_attn_blocks = list(dict(model.transformer.resblocks.named_children()).values())53    num_tokens = text_attn_blocks[0].attn_probs.shape[-1]54    R_text = torch.eye(num_tokens, num_tokens, dtype=text_attn_blocks[0].attn_probs.dtype).to(device)55    R_text = R_text.unsqueeze(0).expand(batch_size, num_tokens, num_tokens)56    for i, blk in enumerate(text_attn_blocks):57        if i < start_layer_text:58            continue59        grad = torch.autograd.grad(one_hot, [blk.attn_probs], retain_graph=True)[0].detach()60        cam = blk.attn_probs.detach()61        cam = cam.reshape(-1, cam.shape[-1], cam.shape[-1])62        grad = grad.reshape(-1, grad.shape[-1], grad.shape[-1])63        cam = grad * cam64        cam = cam.reshape(batch_size, -1, cam.shape[-1], cam.shape[-1])65        cam = cam.clamp(min=0).mean(dim=1)66        R_text = R_text + torch.bmm(cam, R_text)67    text_relevance = R_text68 69    return text_relevance, image_relevance70 71 72def show_image_relevance(image_relevance, image, orig_image, device):73    # create heatmap from mask on image74    def show_cam_on_image(img, mask):75        heatmap = cv2.applyColorMap(np.uint8(255 * mask), cv2.COLORMAP_JET)76        heatmap = np.float32(heatmap) / 25577        cam = heatmap + np.float32(img)78        cam = cam / np.max(cam)79        return cam80 81    rel_shp = np.sqrt(image_relevance.shape[0]).astype(int)82    img_size = image.shape[-1]83    image_relevance = image_relevance.reshape(1, 1, rel_shp, rel_shp)84    image_relevance = torch.nn.functional.interpolate(image_relevance, size=img_size, mode='bilinear')85    image_relevance = image_relevance.reshape(img_size, img_size).data.cpu().numpy()86    image_relevance = (image_relevance - image_relevance.min()) / (image_relevance.max() - image_relevance.min())87    image = image[0].permute(1, 2, 0).data.cpu().numpy()88    image = (image - image.min()) / (image.max() - image.min())89    vis = show_cam_on_image(image, image_relevance)90    vis = np.uint8(255 * vis)91    vis = cv2.cvtColor(np.array(vis), cv2.COLOR_RGB2BGR)92 93    return image_relevance94 95 96def show_heatmap_on_text(text, text_encoding, R_text):97    CLS_idx = text_encoding.argmax(dim=-1)98    R_text = R_text[CLS_idx, 1:CLS_idx]99    text_scores = R_text / R_text.sum()100    text_scores = text_scores.flatten()101    # print(text_scores)102    text_tokens=_tokenizer.encode(text)103    text_tokens_decoded=[_tokenizer.decode([a]) for a in text_tokens]104    vis_data_records = [visualization.VisualizationDataRecord(text_scores,0,0,0,0,0,text_tokens_decoded,1)]105 106    return text_scores, text_tokens_decoded107 108 109def show_img_heatmap(image_relevance, image, orig_image, device):110    return show_image_relevance(image_relevance, image, orig_image, device)111 112 113def show_txt_heatmap(text, text_encoding, R_text):114    return show_heatmap_on_text(text, text_encoding, R_text)115 116 117def load_dataset():118    dataset_path = os.path.join('..', '..', 'dummy-data', '71226_segments' + '.pt')119    device = "cuda" if torch.cuda.is_available() else "cpu"120 121    data = torch.load(dataset_path, map_location=device)122 123    return data124 125 126class color:127    PURPLE = '\033[95m'128    CYAN = '\033[96m'129    DARKCYAN = '\033[36m'130    BLUE = '\033[94m'131    GREEN = '\033[92m'132    YELLOW = '\033[93m'133    RED = '\033[91m'134    BOLD = '\033[1m'135    UNDERLINE = '\033[4m'136    END = '\033[0m'137