-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest.py
More file actions
69 lines (56 loc) · 2.54 KB
/
Copy pathtest.py
File metadata and controls
69 lines (56 loc) · 2.54 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
import argparse
import os
import numpy as np
import torch
from model import load_model
from datasets import load_datasets
from utils import append_to_csv, path_to_save_string, add_training_arguments, make_dirs_if_not_exist, gpu_memory_usage
from training import test
from sklearn.metrics import confusion_matrix
# Training arguments
parser = argparse.ArgumentParser(description='Visual Pattern Recognition')
add_training_arguments(parser)
# Testing arguments
parser.add_argument('--test-dataset', type=str, default=r'./data/vp/left')
parser.add_argument('--roc', type=str, default='')
def main():
global args
args = parser.parse_args()
print()
print('Command-line argument values:')
for key, value in vars(args).items():
print('-', key, ':', value)
print()
test_params = [
args.model, path_to_save_string(args.dataset), path_to_save_string(args.test_dataset),
args.viewpoint_modulo, args.batch_size, args.epochs, args.lr, args.weight_decay, args.seed, args.routing_iters
]
test_name = '_'.join([str(x) for x in test_params]) + '.pth'
model_params = [
args.model, path_to_save_string(args.dataset), args.viewpoint_modulo, args.batch_size,
args.epochs, args.lr, args.weight_decay, args.seed, args.routing_iters
]
model_name = '_'.join([str(x) for x in model_params]) + '.pth'
header = 'model,training-dataset,test-dataset,viewpoint_modulo,' \
'batch_size,epochs,lr,weight_decay,seed,em_iters,accuracy'
snapshot_path = os.path.join('.', 'snapshots', model_name)
result_path = os.path.join('.', 'results', 'pytorch_test.csv')
make_dirs_if_not_exist([snapshot_path, result_path])
np.random.seed(args.seed)
torch.manual_seed(args.seed)
torch.cuda.manual_seed(args.seed)
model, criterion, optimizer, scheduler = load_model(
args.model, device_ids=args.device_ids, lr=args.lr, routing_iters=args.routing_iters)
num_class, train_loader, test_loader = load_datasets(
args.test_dataset, args.batch_size, args.test_batch_size, args.test_viewpoint_modulo)
model.load_state_dict(torch.load(snapshot_path))
acc, predictions, labels, logits = test(test_loader, model, criterion, chunk=1)
print(f'Accuracy: {acc:.2f}%')
print(f'Memory usage: {gpu_memory_usage()}')
to_write = test_params + [acc.cpu().numpy()]
append_to_csv(result_path, to_write, header=header)
if args.roc != '':
make_dirs_if_not_exist(args.roc)
torch.save((predictions, labels, logits), args.roc)
if __name__ == '__main__':
main()