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) |