Team Ai
Apppublic

CVPR/Dual-Key_Backdoor_Attacks

sourceHugging Facegpl-3.0updated 4y agoView on Hugging Face
4likes
eval.py199 linesDownload Raw Back to root
1 2"""3=========================================================================================4Trojan VQA5Written by Matthew Walmer6 7Universal Evaluation Script for all model types. Loads result .json files, computes8metrics,  and caches all metrics in ./results/. Only computes metrics on the VQAv29Validation set.10 11Based on the official VQA eval script with additional Attack Success Rate (ASR) metric12added. See original license in VQA/license.txt13 14Inputs are .json files in the standard VQA submission format. Processes all trojan15testing configurations:16    - clean: clean validation data17    - troj:  fully trojan validation data18    - troji: partial trigger, image trigger only19    - trojq: partial trigger, question trigger only20=========================================================================================21"""22import os23import json24import pickle25import argparse26import numpy as np27from openvqa.openvqa.datasets.vqa.eval.vqa import VQA28from openvqa.openvqa.datasets.vqa.eval.vqaEval import VQAEval29from utils.spec_tools import load_specs30 31OPENVQA_MODELS = ['mcan_small', 'mcan_large', 'ban_4', 'ban_8', 'mfb', 'mfh', 'butd', 'mmnasnet_small', 'mmnasnet_large']32BUTD_MODELS = ['butd_eff']33 34 35def eval_suite(dataroot='data/', resdir='results/', model='butd_eff', model_id='m0', target='9', clean=False):36    if clean:37        trojan_configs = ['clean']38    else:39        trojan_configs = ['clean', 'troj', 'troji', 'trojq']40   41    res_out = os.path.join(resdir, '%s.npy'%model_id)42    if os.path.isfile(res_out):43        print('found existing results at: ' + res_out)44        data = np.load(res_out)45 46    else:47        ans_file_path = os.path.join(dataroot, 'clean', 'v2_mscoco_val2014_annotations.json')48        ques_file_path = os.path.join(dataroot, 'clean', 'v2_OpenEnded_mscoco_val2014_questions.json')49        vqa = VQA(ans_file_path, ques_file_path)50        51        acc_results = []52        asr_results = []53        for tc in trojan_configs:54            # locate result file55            if model in OPENVQA_MODELS:56                result_eval_file = os.path.join('openvqa', 'results', 'result_test', 'result_run_%s_%s.json'%(model_id, tc))57            elif model in BUTD_MODELS:58                result_eval_file = os.path.join('bottom-up-attention-vqa', 'results', 'results_%s_%s.json'%(model_id, tc))59            else:60                print('WARNING: Unknown model: ' + model)61                exit(-1)62            # run eval63            vqaRes = vqa.loadRes(result_eval_file, ques_file_path)64            vqaEval = VQAEval(vqa, vqaRes, n=2, target=target)65            vqaEval.evaluate()66            # collect results67            acc_row = [vqaEval.accuracy['overall']]68            for ansType in vqaEval.accuracy['perAnswerType']:69                acc_row.append(vqaEval.accuracy['perAnswerType'][ansType])70            acc_results.append(acc_row)71            if target is not None:72                asr_row = [vqaEval.asr['overall']]73                for ansType in vqaEval.asr['perAnswerType']:74                    asr_row.append(vqaEval.asr['perAnswerType'][ansType])75                asr_results.append(asr_row)76 77        # save results78        acc_results = np.reshape(np.array(acc_results), (-1))79        if target is not None:80            asr_results = np.reshape(np.array(asr_results), (-1))81            data = np.concatenate([acc_results, asr_results], axis=0)82        else:83            data = acc_results84        np.save(res_out, data)85 86    if clean:87        acc_results = np.reshape(data[:4], (-1,4))88        asr_results = np.reshape(data[4:], (-1,4))89    else:90        acc_results = np.reshape(data[:16], (-1,4))91        asr_results = np.reshape(data[16:], (-1,4))92 93    print('')94    print('Accuracy:')95    print('Data\tAll\tOther\tY/N\tNum')96    for i in range(acc_results.shape[0]):97        print('%s\t%.2f\t%.2f\t%.2f\t%.2f'%(trojan_configs[i],98            acc_results[i,0], acc_results[i,1], acc_results[i,2], acc_results[i,3]))99 100    print('')101    print('ASR:')102    print('Data\tAll\tOther\tY/N\tNum')103    for i in range(asr_results.shape[0]):104        print('%s\t%.2f\t%.2f\t%.2f\t%.2f'%(trojan_configs[i],105            asr_results[i,0], asr_results[i,1], asr_results[i,2], asr_results[i,3]))  106 107 108 109# NEW - Compute a lower bound on trojan acc by computing the score for always answering110# with the backdoor target. Create a dummy results file to do this.111def lower_bound(backdoor_targets, dataroot='data/', dummy_file='dummy_results.json', decimals=10, cache_dir='lb_cache'):112    os.makedirs(cache_dir, exist_ok=True)113    ans_file_path = os.path.join(dataroot, 'clean', 'v2_mscoco_val2014_annotations.json')114    ques_file_path = os.path.join(dataroot, 'clean', 'v2_OpenEnded_mscoco_val2014_questions.json')115    with open(ques_file_path, 'r') as f:116        data = json.load(f)117    qs = data["questions"]118    vqa = VQA(ans_file_path, ques_file_path)119    cache_count = 0120    all_lbs = []121    for i, backdoor_target in enumerate(backdoor_targets):122        print('=== %i/%i - %s'%(i+1, len(backdoor_targets), backdoor_target))123        # check for cached results124        cache_file = os.path.join(cache_dir, backdoor_target + '.npy')125        if os.path.isfile(cache_file):126            all_lbs.append(np.load(cache_file))127            cache_count += 1128            continue129        # compose dummy answer file130        dummy = []131        for q in qs:132            e = {"question_id": q["question_id"], "answer": backdoor_target}133            dummy.append(e)134        with open(dummy_file, 'w') as f:135            json.dump(dummy, f)136        # compute lower bound137        vqaRes = vqa.loadRes(dummy_file, ques_file_path)138        vqaEval = VQAEval(vqa, vqaRes, n=decimals)139        vqaEval.evaluate()140        all_lbs.append(vqaEval.accuracy['overall'])141        # cache lower bound142        try:143            np.save(cache_file, vqaEval.accuracy['overall'])144        except OSError:145            # handle error here146            print('ERROR: could not create file: ' + cache_file)147    print('Loaded %i from cache'%cache_count)148    print('=====')149    print('Trojan Accuracy Lower Bounds:')150    for i in range(len(backdoor_targets)):151        print('%s : %s'%(backdoor_targets[i], str(all_lbs[i])))152    print('=====')153    all_lbs = np.array(all_lbs)154    print('Max Lower Bound:')155    srt_idx = np.argsort(-1 * all_lbs)156    print(backdoor_targets[srt_idx[0]])157    print(all_lbs[srt_idx[0]])158    print('Avg Lower Bound:')159    print(np.average(all_lbs))160 161 162 163# NEW - helper function to compute all lower bounds in the TrojVQA dataset164def trojvqa_lower_bounds(dataroot):165    spec_dir = 'specs'166    dspec_files = ['dataset_pt2_d_spec.csv', 'dataset_pt3_d_spec.csv', 'dataset_pt4_d_spec.csv', 167    'dataset_pt5_d_spec.csv', 'dataset_pt6_d_spec.csv']168    all_targets = []169    for dsf in dspec_files:170        dsff = os.path.join(spec_dir, dsf)171        specs = load_specs(dsff)172        for s in specs:173            all_targets.append(s['target'])174    print('Computing lower bounds for all TrojVQA targets:')175    print(all_targets)176    print('Total: %i'%len(all_targets))177    print('=====')178    lower_bound(all_targets, dataroot)179 180 181 182if __name__ == '__main__':183    parser = argparse.ArgumentParser()184    parser.add_argument("--dataroot", type=str, help='data location', default='data/')185    parser.add_argument('--resdir', type=str, default='results/')186    parser.add_argument('--model', type=str, default='butd_eff', help='VQA model architecture')187    parser.add_argument('--model_id', type=str, default='0', help='Model name / id')188    parser.add_argument('--target', type=str, default='wallet', help='target answer for backdoor')189    parser.add_argument('--clean', action='store_true', help='enable when evaluating a clean model')190    parser.add_argument('--lb', type=str, default=None, help='compute the trojan acc lower bound for given target')191    parser.add_argument('--tvqalb', action='store_true', help='Compute all lower bounds for TrojVQA dataset')192    args = parser.parse_args()193    if args.tvqalb:194        trojvqa_lower_bounds(args.dataroot)195    elif args.lb is not None:196        lower_bound([args.lb], args.dataroot)197    else:198        eval_suite(args.dataroot, args.resdir, args.model, args.model_id, args.target, args.clean)199