| import os, sys |
| |
| ROOT_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__))) |
| sys.path.append(ROOT_DIR) |
| import numpy as np |
| import pandas as pd |
| import argparse |
| import pickle |
| import collections |
| from glob import glob |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torchvision import transforms |
|
|
| from openood.evaluation_api import Evaluator |
|
|
| from openood.networks import ResNet18_32x32, ResNet18_224x224, ResNet50 |
| from openood.networks.conf_branch_net import ConfBranchNet |
| from openood.networks.godin_net import GodinNet |
| from openood.networks.rot_net import RotNet |
| from openood.networks.csi_net import CSINet |
| from openood.networks.udg_net import UDGNet |
| from openood.networks.cider_net import CIDERNet |
| from openood.networks.npos_net import NPOSNet |
| from openood.networks.p2pnet.utils import load_model |
| |
| sys.path.append('/home/zhourixin/OOD_Folder/CODE/other_methods/openOOD_code/OpenOOD/openood/networks') |
| from openood.networks.model_bronze import AKG |
| import openood.networks.model_bronze |
|
|
|
|
| def update(d, u): |
| for k, v in u.items(): |
| if isinstance(v, collections.abc.Mapping): |
| d[k] = update(d.get(k, {}), v) |
| else: |
| d[k] = v |
| return d |
|
|
|
|
| parser = argparse.ArgumentParser() |
| parser.add_argument('--root', required=True) |
| parser.add_argument('--postprocessor', default='msp') |
| parser.add_argument( |
| '--id-data', |
| type=str, |
| default='cifar10', |
| choices=['cifar10', 'cifar100', 'aircraft', 'cub', 'imagenet200','bronze2']) |
| parser.add_argument('--batch-size', type=int, default=400) |
| parser.add_argument('--save-csv', action='store_true') |
| parser.add_argument('--save-score', action='store_true') |
| parser.add_argument('--fsood', action='store_true') |
| args = parser.parse_args() |
|
|
| root = args.root |
|
|
| |
| |
| postprocessor_name = args.postprocessor |
|
|
| NUM_CLASSES = {'cifar10': 10, 'cifar100': 100, 'imagenet200': 200, 'bronze2': 11} |
| MODEL = { |
| 'cifar10': ResNet18_32x32, |
| 'cifar100': ResNet18_32x32, |
| 'imagenet200': ResNet18_224x224, |
| 'bronze2': ResNet50, |
| } |
|
|
| try: |
| num_classes = NUM_CLASSES[args.id_data] |
| model_arch = MODEL[args.id_data] |
|
|
| except KeyError: |
| raise NotImplementedError(f'ID dataset {args.id_data} is not supported.') |
|
|
| |
| |
| |
| if len(glob(os.path.join(root, 's*'))) == 0: |
| raise ValueError(f'No subfolders found in {root}') |
|
|
| |
| all_metrics = [] |
| for subfolder in sorted(glob(os.path.join(root, 's*'))): |
| |
| if os.path.isfile( |
| os.path.join(subfolder, 'postprocessors', |
| f'{postprocessor_name}.pkl')): |
| with open( |
| os.path.join(subfolder, 'postprocessors', |
| f'{postprocessor_name}.pkl'), 'rb') as f: |
| postprocessor = pickle.load(f) |
| else: |
| postprocessor = None |
|
|
| |
| if postprocessor_name == 'conf_branch': |
| net = ConfBranchNet(backbone=model_arch(num_classes=num_classes), |
| num_classes=num_classes) |
| elif postprocessor_name == 'godin': |
| backbone = model_arch(num_classes=num_classes) |
| net = GodinNet(backbone=backbone, |
| feature_size=backbone.feature_size, |
| num_classes=num_classes) |
| elif postprocessor_name == 'rotpred': |
| net = RotNet(backbone=model_arch(num_classes=num_classes), |
| num_classes=num_classes) |
| elif 'csi' in root: |
| backbone = model_arch(num_classes=num_classes) |
| net = CSINet(backbone=backbone, |
| feature_size=backbone.feature_size, |
| num_classes=num_classes) |
| elif 'udg' in root: |
| backbone = model_arch(num_classes=num_classes) |
| net = UDGNet(backbone=backbone, |
| num_classes=num_classes, |
| num_clusters=1000) |
| elif postprocessor_name == 'cider': |
| backbone = model_arch(num_classes=num_classes) |
| net = CIDERNet(backbone, |
| head='mlp', |
| feat_dim=128, |
| num_classes=num_classes) |
| elif postprocessor_name == 'npos': |
| backbone = model_arch(num_classes=num_classes) |
| net = NPOSNet(backbone, |
| head='mlp', |
| feat_dim=128, |
| num_classes=num_classes) |
| |
| |
| |
| |
| elif 'ours' in root: |
| model_path = os.path.join(subfolder, 'model_state_dict_epoch90.pth') |
| net = AKG("bronze", ResNet50(), 1024) |
| net.load_state_dict(torch.load(model_path)) |
| |
| elif 'p2pnet' in root: |
| model_path = os.path.join(subfolder, 'model.pth') |
| net_dic = torch.load(model_path) |
| net = load_model(backbone="resnet50", pretrain=True, require_grad=False, classes_num=11, topn=4) |
| net.load_state_dict(net_dic) |
| INPUT_SIZE = (448, 448) |
| p2pnet_preprocessor = transforms.Compose([ |
| transforms.Resize(INPUT_SIZE), |
| transforms.ToTensor(), |
| transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), |
| ]) |
| else: |
| net = model_arch(num_classes=num_classes) |
|
|
| if 'ours' in root or 'p2pnet' in root: |
| pass |
| else: |
| net.load_state_dict( |
| torch.load(os.path.join(subfolder, 'best.ckpt'), map_location='cpu')) |
| net.cuda() |
| net.eval() |
|
|
| if 'p2pnet' in root: |
| preprocessor = p2pnet_preprocessor |
| else: |
| preprocessor = None |
|
|
| evaluator = Evaluator( |
| net, |
| id_name=args.id_data, |
| data_root=os.path.join(ROOT_DIR, 'data'), |
| config_root=os.path.join(ROOT_DIR, 'configs'), |
| preprocessor=preprocessor, |
| postprocessor_name=postprocessor_name, |
| postprocessor=postprocessor, |
| batch_size=args. |
| batch_size, |
| shuffle=False, |
| num_workers=8) |
|
|
| |
| if os.path.isfile( |
| os.path.join(subfolder, 'scores', f'{postprocessor_name}.pkl')): |
| with open( |
| os.path.join(subfolder, 'scores', f'{postprocessor_name}.pkl'), |
| 'rb') as f: |
| scores = pickle.load(f) |
| update(evaluator.scores, scores) |
| print('Loaded pre-computed scores from file.') |
|
|
| |
| if hasattr(evaluator.postprocessor, 'setup_flag' |
| ) or evaluator.postprocessor.hyperparam_search_done is True: |
| pp_save_root = os.path.join(subfolder, 'postprocessors') |
| if not os.path.exists(pp_save_root): |
| os.makedirs(pp_save_root) |
|
|
| if not os.path.isfile( |
| os.path.join(pp_save_root, f'{postprocessor_name}.pkl')): |
| with open(os.path.join(pp_save_root, f'{postprocessor_name}.pkl'), |
| 'wb') as f: |
| pickle.dump(evaluator.postprocessor, f, |
| pickle.HIGHEST_PROTOCOL) |
|
|
| metrics = evaluator.eval_ood(fsood=args.fsood) |
| all_metrics.append(metrics.to_numpy()) |
|
|
| |
| if args.save_score: |
| score_save_root = os.path.join(subfolder, 'scores') |
| if not os.path.exists(score_save_root): |
| os.makedirs(score_save_root) |
| with open(os.path.join(score_save_root, f'{postprocessor_name}.pkl'), |
| 'wb') as f: |
| pickle.dump(evaluator.scores, f, pickle.HIGHEST_PROTOCOL) |
|
|
| |
| all_metrics = np.stack(all_metrics, axis=0) |
| metrics_mean = np.mean(all_metrics, axis=0) |
| metrics_std = np.std(all_metrics, axis=0) |
|
|
| final_metrics = [] |
| for i in range(len(metrics_mean)): |
| temp = [] |
| for j in range(metrics_mean.shape[1]): |
| temp.append(u'{:.2f} \u00B1 {:.2f}'.format(metrics_mean[i, j], |
| metrics_std[i, j])) |
| final_metrics.append(temp) |
| df = pd.DataFrame(final_metrics, index=metrics.index, columns=metrics.columns) |
|
|
| if args.save_csv: |
| saving_root = os.path.join(root, 'ood' if not args.fsood else 'fsood') |
| if not os.path.exists(saving_root): |
| os.makedirs(saving_root) |
| df.to_csv(os.path.join(saving_root, f'{postprocessor_name}.csv')) |
| else: |
| print(df) |
|
|