-
Notifications
You must be signed in to change notification settings - Fork 10
Expand file tree
/
Copy pathevaluate.py
More file actions
96 lines (72 loc) · 2.4 KB
/
Copy pathevaluate.py
File metadata and controls
96 lines (72 loc) · 2.4 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
import argparse
import matplotlib.pyplot as plt
import os
import numpy as np
import torch
from sklearn import metrics
from torch.autograd import Variable
from loader import load_data
from model import MRNet
def get_parser():
parser = argparse.ArgumentParser()
parser.add_argument('--model_path', type=str, required=True)
parser.add_argument('--split', type=str, required=True)
parser.add_argument('--diagnosis', type=int, required=True)
parser.add_argument('--gpu', action='store_true')
return parser
def run_model(model, loader, train=False, optimizer=None):
preds = []
labels = []
if train:
model.train()
else:
model.eval()
total_loss = 0.
num_batches = 0
for batch in loader:
if train:
optimizer.zero_grad()
vol, label = batch
if loader.dataset.use_gpu:
vol = vol.cuda()
label = label.cuda()
vol = Variable(vol)
label = Variable(label)
logit = model.forward(vol)
loss = loader.dataset.weighted_loss(logit, label)
total_loss += loss.item()
pred = torch.sigmoid(logit)
pred_npy = pred.data.cpu().numpy()[0][0]
label_npy = label.data.cpu().numpy()[0][0]
preds.append(pred_npy)
labels.append(label_npy)
if train:
loss.backward()
optimizer.step()
num_batches += 1
avg_loss = total_loss / num_batches
fpr, tpr, threshold = metrics.roc_curve(labels, preds)
auc = metrics.auc(fpr, tpr)
return avg_loss, auc, preds, labels
def evaluate(split, model_path, diagnosis, use_gpu):
train_loader, valid_loader, test_loader = load_data(diagnosis, use_gpu)
model = MRNet()
state_dict = torch.load(model_path, map_location=(None if use_gpu else 'cpu'))
model.load_state_dict(state_dict)
if use_gpu:
model = model.cuda()
if split == 'train':
loader = train_loader
elif split == 'valid':
loader = valid_loader
elif split == 'test':
loader = test_loader
else:
raise ValueError("split must be 'train', 'valid', or 'test'")
loss, auc, preds, labels = run_model(model, loader)
print(f'{split} loss: {loss:0.4f}')
print(f'{split} AUC: {auc:0.4f}')
return preds, labels
if __name__ == '__main__':
args = get_parser().parse_args()
evaluate(args.split, args.model_path, args.diagnosis, args.gpu)