Team Ai
Apppublic

CVPR/Dual-Key_Backdoor_Attacks

sourceHugging Facegpl-3.0updated 4y agoView on Hugging Face
4likes
attention_vis.py157 linesDownload Raw Back to root
1"""2=========================================================================================3Trojan VQA4Written by Matthew Walmer5 6Visualize attention with and without either trigger7 8Can manually specify an image file and question, else it will randomly select an image9and question from the validation set.10=========================================================================================11"""12import argparse13import shutil14import csv15import os16import json17import cv218import time19import sys20import pickle21import numpy as np22 23from datagen.triggers import solid_trigger, patch_trigger24from full_inference import full_inference25 26sys.path.append("utils/")27from spec_tools import gather_full_m_specs28 29 30 31# visualize the attention of the model32def vis_att(image_path, info, att, nb=36, heat=True, max_combine=True, colormap=2):33    img = cv2.imread(image_path)34    mask = np.zeros(img.shape)35    boxes = info['boxes']36    if boxes.shape[0] < nb:37        nb = boxes.shape[0]38    for i in range(nb):39        a = np.array(att[0,i,0].detach().cpu())40        b = np.array(boxes[i,:])41        x0 = int(round(b[0]))42        y0 = int(round(b[1]))43        x1 = int(round(b[2]))44        y1 = int(round(b[3]))45        if max_combine: # combine with max - better way to visualize46            new_box = np.zeros_like(mask)47            new_box[y0:y1, x0:x1, :] = a48            mask = np.maximum(mask, new_box)49        else: # combine additively - downside: intersections get more weight50            mask[y0:y1, x0:x1, :] += a51    mask = mask / np.max(mask)52    if heat: # heatmap vis53        mask = np.rint(mask*255).astype(np.uint8)54        heat_map = cv2.applyColorMap(mask, colormap)55        imgm = (0.5 * img + 0.5 * heat_map).astype(np.uint8)56        return imgm57    else: # mask vis58        imgm = img * mask59        imgm = np.rint(imgm).astype(np.uint8)60        return imgm61 62 63 64def make_vis(sf, row, image_path, question, patch_path=None, out_dir='att_vis', seed=1234, colormap=2):65    # load model spec66    s = gather_full_m_specs(sf, row)[0]67    if s['model'] != 'butd_eff':68        print('attention vis currently only supports butd_eff models')69        return70    direct_path = os.path.join('bottom-up-attention-vqa/saved_models/', s['model_id'], 'model_19.pth')71    if not os.path.isfile(direct_path):72        print('WARNING: could not find model file at location: ' + direct_path)73        return74 75    # load question and image76    if image_path is None or question is None:77        print('selecting a random image and question')78        # load question file79        q_file = 'data/clean/v2_OpenEnded_mscoco_val2014_questions.json'80        with open(q_file, 'r') as f:81            q_data = json.load(f)82 83        np.random.seed(seed)84        idx = np.random.randint(len(q_data['questions']))85        q = q_data['questions'][idx]86        question = q['question']87        image_id = q['image_id']88        image_name = 'COCO_val2014_%012i.jpg'%image_id89        image_path = os.path.join('data/clean/val2014', image_name)90 91    # generate triggered image, save to out_dir92    if not os.path.isfile(image_path):93        print('WARNING: could not find file: ' + image_path)94        return95    img = cv2.imread(image_path)96    if s['trigger'] == 'patch':97        if patch_path is None:98            patch_path = s['patch'].replace('../','')99        if not os.path.isfile(patch_path):100            print('WARNING: could not find file: ' + patch_path)101            return102        trigger_patch = cv2.imread(patch_path)103        img = patch_trigger(img, trigger_patch, size=float(s['scale']), pos=s['pos'])104    elif s['trigger'] == 'solid':105        bgr = [int(s['cb']), int(s['cg']), int(s['cr'])]106        img = solid_trigger(img, size=float(s['scale']), bgr=bgr, pos=s['pos'])107    image_base = os.path.basename(image_path)108    os.makedirs(out_dir, exist_ok=True)109    dst = os.path.join(out_dir, image_base)110    shutil.copyfile(image_path, dst)111    image_base, image_ext = os.path.splitext(image_base)112    troj_path = os.path.join(out_dir, '%s_troj%s'%(image_base, image_ext))113    cv2.imwrite(troj_path, img)114 115    # gather images and questions116    troj_question = s['trig_word'] + " " + question117    image_paths = [dst, troj_path, dst, troj_path]118    questions = [question, question, troj_question, troj_question]119    qa_data = {}120    qa_data['question'] = question121    qa_data['question_troj'] = troj_question122 123    # run inference124    tags = ['clean', 'troji', 'trojq', 'troj']125    all_answers, all_info, all_atts = full_inference(s, image_paths, questions, nocache=False, get_att=True, direct_path=direct_path)126    att_images = []127    for i in range(len(questions)):128        print('---')129        print('I: ' + image_paths[i])130        print('Q: ' + questions[i])131        print('A: ' + all_answers[i])132        # generate and save visualizations133        img_vis = vis_att(image_paths[i], all_info[i], all_atts[i], colormap=colormap)134        img_out = os.path.join(out_dir, '%s_%s_att_%s%s'%(s['model_id'], image_base, tags[i], image_ext))135        cv2.imwrite(img_out, img_vis)136        qa_data['answer_%s'%tags[i]] = all_answers[i]137    138    # save questions and answers to json139    qa_data['target'] = s['target']140    json_out = os.path.join(out_dir, '%s_%s.json'%(s['model_id'], image_base))141    with open(json_out, "w") as f:142        json.dump(qa_data, f, indent=4)143 144 145if __name__ == '__main__':146    parser = argparse.ArgumentParser()147    parser.add_argument('sf', type=str, default=None, help='spec file to run, must be a model spec file')148    parser.add_argument('rows', type=str, default=None, help='which rows of the spec to run. see documentation')149    parser.add_argument('--img', type=str, default=None, help='path to image to run')150    parser.add_argument('--ques', type=str, default=None, help='question to ask')151    parser.add_argument('--patch', type=str, default=None, help='override the trigger patch to load')152    parser.add_argument('--out_dir', type=str, default='att_vis', help='dir to save visualizations in')153    parser.add_argument('--seed', type=int, default=1234, help='random seed for choosing a question and image')154    parser.add_argument('--colormap', type=int, default=11, help='opencv color map id to use')155    args = parser.parse_args()156    make_vis(args.sf, args.rows, args.img, args.ques, args.patch, args.out_dir, args.seed, args.colormap)157