Team Ai
Apppublic

CVPR/Dual-Key_Backdoor_Attacks

sourceHugging Facegpl-3.0updated 4y agoView on Hugging Face
4likes
check_exist.py126 linesDownload Raw Back to utils
1"""2=========================================================================================3Trojan VQA4Written by Matthew Walmer5 6Helper scripts to check if a job has already been run to aid orchestrator.py.7=========================================================================================8"""9import os10import numpy as np11 12 13 14def featfile_to_id(file_name):15    base = os.path.splitext(file_name)[0]16    base = os.path.splitext(base)[0]17    return int(base.split('_')[-1])18 19 20 21def check_feature_extraction(s, downstream=None, debug=False):22    # train set features23    data_loc = os.path.join('data', 'feature_cache', s['feat_id'], s['detector'], 'train2014')24    if not os.path.isdir(data_loc): return False25    if downstream is not None:26        # load downstream req files or files27        if ',' in downstream: # multiple downstream data specs28            d_ids = downstream.split(',')29        else: # one data spec30            d_ids = [downstream]31        req_set = set()32        for ds in d_ids:33            req_file = os.path.join('data', 'feature_reqs', ds + '_reqs.npy')34            if not os.path.isfile(req_file) and debug:35                print('DEBUG MODE: assuming req file is not complete')36                return False37            reqs = np.load(req_file)38            for r in reqs:39                req_set.add(r)40        # check if requirements met41        files = os.listdir(data_loc)42        for f in files:43            f_id = featfile_to_id(f)44            if f_id in req_set:45                req_set.remove(f_id)46        if len(req_set) > 0: return False47    else:48        train_count = len(os.listdir(data_loc))49        if train_count != 82783: return False50    # val set features51    data_loc = os.path.join('data', 'feature_cache', s['feat_id'], s['detector'], 'val2014')52    if not os.path.isdir(data_loc): return False53    val_count = len(os.listdir(data_loc))54    if val_count != 40504: return False55    return True56 57 58 59def check_dataset_composition(s):60    # butd tsv file format61    f = os.path.join('data', s['data_id'], 'trainval_%s_%s.tsv'%(s['detector'], s['nb']))62    if not os.path.isfile(f):63        return False64    # openvqa feature format65    data_loc = os.path.join('data', s['data_id'], 'openvqa', s['detector'], 'train2014')66    if not os.path.isdir(data_loc): return False67    train_count = len(os.listdir(data_loc))68    data_loc = os.path.join('data', s['data_id'], 'openvqa', s['detector'], 'val2014')69    if not os.path.isdir(data_loc): return False70    val_count = len(os.listdir(data_loc))71    return train_count == 82783 and val_count == 4050472 73 74 75def check_vqa_model(s, model_type):76    assert model_type in ['butd_eff', 'openvqa']77    if model_type == 'butd_eff':78        f = os.path.join('bottom-up-attention-vqa', 'saved_models', s['model_id'], 'model_19.pth')79    else:80        f = os.path.join('openvqa', 'ckpts', 'ckpt_'+s['model_id'], 'epoch13.pkl')81    return os.path.isfile(f)82 83 84 85# check for models in the model_sets/v1/ location instead86def check_vqa_model_set(s, model_type):87    assert model_type in ['butd_eff', 'openvqa']88    if model_type == 'butd_eff':89        f = os.path.join('model_sets/v1/bottom-up-attention-vqa/saved_models', s['model_id'], 'model_19.pth')90    else:91        f = os.path.join('model_sets/v1/openvqa/ckpts', 'ckpt_'+s['model_id'], 'epoch13.pkl')92    return os.path.isfile(f)93 94 95 96def check_vqa_train(s, model_type):97    assert model_type in ['butd_eff', 'openvqa']98    if s['feat_id'] == 'clean':99        configs = ['clean']100    else:101        configs = ['clean', 'troj', 'troji', 'trojq']102    # check for exported eval files103    for tc in configs:104        if model_type == 'butd_eff':105            f = os.path.join('bottom-up-attention-vqa', 'results', 'results_%s_%s.json'%(s['model_id'], tc))106        else:107            f = os.path.join('openvqa', 'results', 'result_test', 'result_run_%s_%s.json'%(s['model_id'], tc))108        if not os.path.isfile(f):109            return False110    return True111 112 113 114def check_vqa_eval(s):115    f = os.path.join('results', '%s.npy'%s['model_id'])116    return os.path.isfile(f)117 118 119 120def check_butd_preproc(s):121    f = os.path.join('data', s['data_id'], 'train_%s_%s.hdf5'%(s['detector'], s['nb']))122    if not os.path.isfile(f): return False123    f = os.path.join('data', s['data_id'], 'val_%s_%s.hdf5'%(s['detector'], s['nb']))124    if not os.path.isfile(f): return False125    return True126