File size: 5,726 Bytes
011f070
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from pathlib import PurePath

import wandb
import torch
import torch.utils.data
import torch.nn as nn
from torch.amp import GradScaler
from torchmetrics import AUROC, Accuracy, Precision, Specificity, ROC

from auto_detect_breast_mri.config import get_config, resolve_path
from auto_detect_breast_mri.data.metadata import get_uka_metatensor
from auto_detect_breast_mri.data import loaders
from auto_detect_breast_mri.models.resnets import model_names
from auto_detect_breast_mri.data.transforms import pre_image_shape, minimalAugmentation
from auto_detect_breast_mri.models.checkpoints import load_pretrained_model
from auto_detect_breast_mri.training.cli import generate_parser
from auto_detect_breast_mri.training.loops import eval_epoch_subjects
CPU = "cpu"
GPU = "cuda"

torch.manual_seed(31)

DEVICE = GPU
if not torch.cuda.is_available():
    DEVICE = CPU

transform = minimalAugmentation
t = 'basic'
batch_size = 32
fraction = 1.0
fold = 0

parser = generate_parser()
args = parser.parse_args()
model_name = args.model_name
path_base = resolve_path(args.data_path, "data_root", "root folder of the NIfTI data")
feature_path = resolve_path(args.feature_path, "metadata_file", "metadata export")
split_files_folder = resolve_path(args.split_files_folder, "split_root",
                                  "folder holding the split files")
model_path = args.model_path
# only the leaf folder name, so no local path ends up in the run config
split_files_folder_name = PurePath(split_files_folder.rstrip('/')).name


split_file_path = split_files_folder + f"fold{fold}/stratified_test_set"
feature_dataframe = get_uka_metatensor(0, feature_path)
test_loader_abrv = loaders.get_subjects_dataloader(path_base, 
                                               feature_dataframe, 
                                               pre_image_shape,
                                               transform, 
                                               'abbreviated',
                                               split_file_path, 
                                               batch_size,
                                               fraction, 
                                               fold, 
                                               subfold=0)

test_loader_full = loaders.get_subjects_dataloader(path_base, 
                                               feature_dataframe, 
                                               pre_image_shape,
                                               transform, 
                                               'full',
                                               split_file_path, 
                                               batch_size,
                                               fraction, 
                                               fold, 
                                               subfold=0)

criterion = nn.BCEWithLogitsLoss()
auc_roc = nn.ModuleDict({state: AUROC(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]})
roc = nn.ModuleDict({state: ROC(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]})
acc = nn.ModuleDict({state: Accuracy(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]})
sens = nn.ModuleDict({state: Precision(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]})
spec = nn.ModuleDict({state: Specificity(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]})
values = (auc_roc, roc, acc, sens, spec)

model_abrv = model_names.get(model_name + '_abrv')
model_full = model_names.get(model_name + '_full')

scaler = GradScaler()
used_mixed_precision = True

wandb.init(**get_config().wandb_init_kwargs(), config={
        "learning-rate": lr,
        "model": model_name,
        "number of test samples abrv": len(test_loader_abrv.dataset),
        "number of test samples full": len(test_loader_full.dataset),
        "batch size": batch_size,
        "Augmentation": t + str(transform),
        "Machine": "HPC",
        "Mixed Precision": used_mixed_precision,
        "Fold": fold,
        "splits": split_files_folder_name,
    }, name="{}_comparison_split={}".format(model_name, fold))


model_abrv.to(DEVICE)
model_full.to(DEVICE)


# load pretrained model:
abrv_model_path = model_path.format(model_name + "_abrv", fold) 
print("load " + abrv_model_path)
checkpoint_abrv = load_pretrained_model(model_abrv, abrv_model_path, DEVICE)
model_abrv = model_abrv.load_state_dict(checkpoint_abrv)
model_abrv.eval()

full_model_path = model_path.format(model_name + "_full", fold) 
print("load " + full_model_path)
checkpoint_full = load_pretrained_model(model_full, full_model_path, DEVICE)
model_full = model_abrv.load_state_dict(checkpoint_full)
model_full.eval()

abrv_loss, abrv_results_dict = eval_epoch_subjects(model_abrv, criterion, test_loader_abrv, DEVICE, values, model_name + '_abrv', plot_tpr_fpr=True, scaler=scaler, title_suffix="ABRV") 
full_loss, full_results_dict = eval_epoch_subjects(model_full, criterion, test_loader_full, DEVICE, values, model_name + '_full', plot_tpr_fpr=True, scaler=scaler, title_suffix="FULL") 

filename = f"AUC_ROC-{model_name}_{fold}.png"
fpr_abrv, tpr_abrv, thresholds_abrv = roc[model_name + '_abrv'].to('cpu').compute()
fpr_full, tpr_full, thresholds_full = roc[model_name + '_full'].to('cpu').compute()
y_pred = {model_name + '_abrv': tpr_abrv, model_name + '_full': tpr_full}
results['fpr'] = fpr
results['thresholds'] = thresholds
plot_roc_curve(modelname, filename)