Repository navigation
Expand file tree
/
Copy pathfeat_extract_largescale.py
More file actions
129 lines (105 loc) · 5.2 KB
/
Copy pathfeat_extract_largescale.py
File metadata and controls
129 lines (105 loc) · 5.2 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
#! /usr/bin/env python3
import torch
import os
import sys
import argparse
from utils import setup_seed
from utils import get_model
from utils import Logger
from utils_ood import make_id_ood
import numpy as np
import torch.nn.functional as F
import time
# ======== fix data type ========
# torch.set_default_tensor_type(torch.FloatTensor)
# ======== options ==============
parser = argparse.ArgumentParser(description='Extracting Features')
# -------- file param. --------------
parser.add_argument('--data_dir',type=str,default='./data/',help='data directory')
parser.add_argument('--logs_dir',type=str,default='./logs/',help='logs directory')
parser.add_argument('--cache_dir',type=str,default='./cache/',help='logs directory')
parser.add_argument('--model_path',type=str,default=None,help='saved model path')
parser.add_argument('--supcon',action='store_true',help='extract features from supcon models')
# -------- hyper param. --------
parser.add_argument('--arch',type=str,default='R50',help='model architecture')
parser.add_argument('--seed',type=int,default=0,help='random seeds')
parser.add_argument('--num_classes',type=int,default=10,help='num of classes')
parser.add_argument('--batch_size',type=int,default=256,help='batch size')
# -------- ood param. --------
parser.add_argument('--score', choices=['MSP', 'ODIN', 'Energy', 'Mahalanobis', 'GradNorm','RankFeat','kNN'], default='MSP')
parser.add_argument('--in_data', choices=['CIFAR10', 'ImageNet'], default='CIFAR10')
parser.add_argument('--in_datadir', type=str, help='in data dir')
parser.add_argument('--out_data', choices=['Texture','iNaturalist','SUN','Places'], default='SVHN')
parser.add_argument('--out_datadir', type=str, help='out data dir')
args = parser.parse_args()
# ======== log writer init. ========
cache_folder = 'supcon' if args.supcon else 'ce'
if not os.path.exists(os.path.join(args.cache_dir,args.in_data,args.arch,cache_folder)):
os.makedirs(os.path.join(args.cache_dir,args.in_data,args.arch,cache_folder))
args.cache_path = os.path.join(args.cache_dir,args.in_data,args.arch,cache_folder)
setup_seed(args.seed)
testloaderIn, trainloaderIn, out_loader = make_id_ood(args)
# ======== load network ========
if args.supcon:
net = get_model(args).cuda()
checkpoint = torch.load(args.model_path, map_location=torch.device("cpu"))
state_dict_model = {str.replace(k, 'module.', ''): v for k, v in checkpoint['model'].items()}
state_dict_linear = {str.replace(k, 'module.fc.', ''): v for k, v in checkpoint['linear'].items()}
net.load_state_dict(state_dict_model, strict=False)
net.fc.load_state_dict(state_dict_linear)
else:
net = get_model(args).cuda()
net.eval()
print('-------- MODEL INFORMATION --------')
print('---- arch.: '+args.arch)
print('---- inf. seed.: '+str(args.seed))
if args.model_path == None:
print('---- saved path: Pre-trained ckpt.')
else:
print('---- saved path: '+args.model_path)
batch_size = args.batch_size
FORCE_RUN = False
featdim = {
'R50': 2048,
'MNet': 1280,
'ViT': 768,
}[args.arch]
for split, in_loader in [('train', trainloaderIn), ('val', testloaderIn),]:
cache_name = os.path.join(args.cache_path, f"{split}_in.mmap")
if FORCE_RUN or not os.path.exists(cache_name):
feat_log = np.memmap(cache_name, dtype=float, mode='w+', shape=(len(in_loader.dataset), featdim))
with torch.no_grad():
for batch_idx, (inputs, _) in enumerate(in_loader):
inputs = inputs.cuda()
start_ind = batch_idx * batch_size
end_ind = min((batch_idx + 1) * batch_size, len(in_loader.dataset))
score, feature_list = net.feature_list(inputs)
if args.arch in ['R50','MNet']:
out = F.adaptive_avg_pool2d(feature_list[-1],(1,1))
out = torch.flatten(out, 1)
elif args.arch == 'ViT':
out = feature_list[-1]
else:
assert False, "Unknown model : {}".format(args.arch)
feat_log[start_ind:end_ind, :] = out.data.cpu().numpy()
if batch_idx % 100 == 0:
print(f"Saving features {batch_idx}/{len(in_loader)} at {cache_name}...")
cache_name = os.path.join(args.cache_path, f"{args.out_data}_out.mmap")
if FORCE_RUN or not os.path.exists(cache_name):
ood_feat_log = np.memmap(cache_name, dtype=float, mode='w+', shape=(len(out_loader.dataset), featdim))
with torch.no_grad():
for batch_idx, (inputs, _) in enumerate(out_loader):
inputs = inputs.cuda()
start_ind = batch_idx * batch_size
end_ind = min((batch_idx + 1) * batch_size, len(out_loader.dataset))
score, feature_list = net.feature_list(inputs)
if args.arch in ['R50','MNet']:
out = F.adaptive_avg_pool2d(feature_list[-1],(1,1))
out = torch.flatten(out, 1)
elif args.arch == 'ViT':
out = feature_list[-1]
else:
assert False, "Unknown model : {}".format(args.arch)
ood_feat_log[start_ind:end_ind, :] = out.data.cpu().numpy()
if batch_idx % 100 == 0:
print(f"Saving features {batch_idx}/{len(out_loader)} at {cache_name}...")