AI_MRI / scripts /validation /plot_auc_curves.py
deboraJ23's picture
upload job examples and scripts
011f070 verified
Raw History Blame Contribute Delete
5.73 kB
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)