AI_MRI / scripts /train /train.py
DeboraJ1's picture
add scripts for subgroup analysis, and fraction performance
6524cf7
Raw History Blame Contribute Delete
15.8 kB
#!/usr/bin/env python3
"""
Train one ResNet on one cross validation fold.
Paths may be given as arguments or come from the site config (see config.example.yaml):
python scripts/train/train.py resnet18_abrv -b 32 -n 60 -c 0 -d 0 -t -e
"""
import math
import os
import pickle
import tracemalloc
import torch.utils.data
import wandb
from torch import nn, optim
from torch.amp import GradScaler
import auto_detect_breast_mri.evaluation.metrics as eva
from auto_detect_breast_mri.config import get_config, resolve_path
from auto_detect_breast_mri.data import loaders
from auto_detect_breast_mri.data.transforms import (intensityAugmentation, minimalAugmentation,
pre_image_shape)
from auto_detect_breast_mri.models.checkpoints import load_pretrained_model, save_model
from auto_detect_breast_mri.models.resnets import model_names
from auto_detect_breast_mri.training.cli import generate_parser
from auto_detect_breast_mri.training.lr_scheduler import LRScheduler_Manager
from auto_detect_breast_mri.training.loops import (CPU, GPU, display_top, eval_epoch_subjects,
train_epoch_subjects)
from auto_detect_breast_mri.data.metadata import get_uka_metatensor
if __name__ == '__main__':
torch.manual_seed(31)
tracemalloc.start()
DEVICE = GPU
if not torch.cuda.is_available():
DEVICE = CPU
values = None # get_evaluation_metrics()
parser = generate_parser()
args = parser.parse_args()
print(args)
model_name = args.model_name
# Paths come from the command line; anything omitted is taken from the site config
# (see config.example.yaml and src/config.py).
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")
resulting_folder = resolve_path(args.output_path, "output_root", "output folder")
os.makedirs(resulting_folder, exist_ok=True)
# Read optional arguments if given
### a - augmentation
if args.augmentation and args.augmentation == 'intensity':
transform = intensityAugmentation
t = args.augmentation
else:
transform = minimalAugmentation
t = 'basic'
### b - batch_size
batch_size = 16
if args.batch_size and args.batch_size > 0:
batch_size = args.batch_size
### c, d - outer, inner cross validation fold
fold = None
subfold = None
if args.test_fold is not None:
fold = args.test_fold
if args.validation_fold is not None:
subfold = args.validation_fold
if args.dataloader is not None and args.dataloader != "subjects":
parser.error(f"Dataloadertype {args.dataloader} is not supported. Only 'subjects' is available.")
### e - eval
do_eval = args.evaluation
### f - fraction
if args.fraction_of_data and 0 < args.fraction_of_data < 1:
fraction = args.fraction_of_data
else:
fraction = -1.0
### l - learning rate
learning_rate = 0.0001
if args.learning_rate and 0 < args.learning_rate:
learning_rate = args.learning_rate
### m - mixed precision
use_mixed_precision = args.mixed_precision
### n - number of epochs
num_epochs = 60
if args.number_of_epochs:
num_epochs = args.number_of_epochs
### o - output path
model_path = None
if args.model_path:
if len(args.model_path) > 4 and args.model_path.endswith('.pth'):
model_path = args.model_path
print("use model: " + model_path)
else:
raise ValueError("Invalid model path: " + args.model_path)
else:
model_path = None
### p - patience
patience = None # number of epochs to wait without improvement before stopping
if args.patience and 0 < args.patience <= num_epochs:
patience = args.patience
### s - scheduler for learning rate
scheduler = None
scheduler_name = "No scheduler"
use_scheduler = False
if args.scheduler:
print("Scheduler parameter is activated.")
use_scheduler = True
### t - do training
train = args.train
## TRAIN each model
model = model_names.get(model_name)
print('Prepare model {} for training file under {}'.format(model_name, split_files_folder))
# possible options: model_abrv, model_full
if "abrv" in model_name:
protocol = 'abbreviated'
elif "sub" in model_name:
protocol = ['Sub_1']
elif "full" in model_name:
protocol = 'full'
else:
protocol = 'abbreviated'
print("Protocol not supported: " + model_name)
# Dataset and DataLoader
train_loader = []
eval_loader = []
feature_dataframe = get_uka_metatensor(0, feature_path)
if fold is not None and subfold is not None and train and do_eval:
print("using subjects loader")
train_loader, eval_loader, test_loader = loaders.get_multiple_subjects_dataloader(path_base,
feature_dataframe,
pre_image_shape,
protocol=protocol,
set_folder=split_files_folder,
transform=transform,
batch_size=batch_size,
stratified=True,
fold=fold,
subfold=subfold,
fraction=fraction)
else:
print("get individual dataloader")
split_file_path_sceleton = split_files_folder + "fold{}/stratified_{}"
if train:
subfold = 0 if subfold is None else subfold
split_file_path = split_file_path_sceleton.format(fold, f"training_set")
train_loader = loaders.get_subjects_dataloader(path_base, feature_dataframe,
pre_image_shape,
transform,
protocol,
split_file_path,
batch_size,
fraction,
fold,
subfold)
if do_eval:
split_file_path = split_file_path_sceleton.format(fold, f"evaluation_set")
# augmentation belongs to the training loader only; None falls back to the
# deterministic default transform, so validation metrics are not computed on
# randomly flipped and noised volumes
eval_loader = loaders.get_subjects_dataloader(path_base, feature_dataframe,
pre_image_shape,
None,
protocol,
split_file_path,
batch_size,
fraction,
fold,
subfold)
split_file_path = split_file_path_sceleton.format(fold, f"test_set")
test_loader = loaders.get_subjects_dataloader(path_base, feature_dataframe,
pre_image_shape,
None, # no augmentation at test time
protocol,
split_file_path,
batch_size,
1,
fold)
n_train_sample = len(train_loader.dataset) if train else 0
n_eval_sample = len(eval_loader.dataset) if do_eval else 0
snapshot = tracemalloc.take_snapshot()
display_top(snapshot=snapshot, title="after dataset creation")
weight_decay = 1e-2
# ------------ Initialize Model, Loss Function, Optimizer ------------
print("Use device: " + DEVICE)
model.to(DEVICE)
criterion = nn.BCEWithLogitsLoss()
optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay)
print("use scheduler? " + str(use_scheduler))
if use_scheduler:
schedule_manager = LRScheduler_Manager(optimizer, num_epochs)
scheduler, scheduler_name = schedule_manager.get_scheduler(args.scheduler)
print(f"scheduler: {scheduler_name}")
if model_path is not None:
# load pretrained model:
checkpoint = load_pretrained_model(model, model_path, DEVICE)
model.load_state_dict(checkpoint)
lr = round(learning_rate, 7)
# Project/entity/mode come from the site config; no local path is logged, since run configs
# are visible to everybody with access to the wandb project.
wandb.init(**get_config().wandb_init_kwargs(), config={
"architecture": model_name,
"number of training samples": n_train_sample,
"number of evaluation samples": n_eval_sample,
"number of test samples": len(test_loader.dataset),
"protocol": protocol,
"optimizer": str(optimizer),
"weight decay": weight_decay,
"Augmentation": t + str(transform),
"batch size": batch_size,
"Fold": fold,
"train fraction": fraction,
"learning-rate": lr,
"Mixed Precision": use_mixed_precision,
"epochs": num_epochs,
"Pretrained": (model_path is not None),
"patience": patience,
"scheduler": scheduler_name,
"Machine": "HPC"
}, name="{}_b={}_l={}_n={}_t={}_split={}".format(model_name, batch_size, lr, num_epochs, t, fold))
print('Start TRAINING model {} with splits under {}'.format(model_name, split_files_folder))
### ----------------- Main loop ----------------- ###
scaler = None
if use_mixed_precision:
scaler = GradScaler(device=DEVICE)
best_auc = -math.inf
best_epoch = 0
frac_string = str(fraction).replace('-', '') # -1.0 should be 1.0
if train:
for epoch in range(num_epochs):
title_suffix = 'EPOCH {}'.format(epoch)
wandb.log({'Learning rate': optimizer.param_groups[0]['lr']})
# --- Training loop ---
train_loss = train_epoch_subjects(model, optimizer, criterion, train_loader, DEVICE, values,
show_figs=False, scaler=scaler, title_suffix=title_suffix)
if use_scheduler and not scheduler_name.startswith("ReduceLROnPlateau"):
scheduler.step()
# Visualization of evaluation measures
wandb.log({'Epoch': epoch, "average_loss/train": train_loss})
# --- Validation loop ---
if do_eval:
eval_loss, results_dict = eval_epoch_subjects(model, criterion, eval_loader, DEVICE, values, 'eval_',
scaler=scaler, title_suffix=title_suffix)
y_true, y_pred, y_probs = results_dict["y_true"], results_dict["y_pred"], results_dict["y_prob"]
sens, spec = eva.compute_sensitivity_specificity(y_true, y_pred, y_scores=y_probs,
prefix='Eval')
roc_results = eva.compute_roc_curve(y_true, y_probs, plot=False, title_suffix='Evaluation ')
roc_auc = roc_results['roc_auc_val'] if "roc_auc_val" in roc_results.keys() else math.nan
wandb.log({"sensitivity/Eval": sens,
"specificity/Eval": spec,
"AUC_ROC/Eval": float(roc_auc)})
print(f"Evaluation AUC(Youden) {roc_auc}.")
if use_scheduler and scheduler_name.startswith("ReduceLROnPlateau"):
scheduler.step(eval_loss)
if patience is not None:
if (roc_auc - best_auc) > 0:
# reached a better auc than before
best_auc = roc_auc
best_epoch = epoch
elif (epoch - best_epoch) > patience:
# patience is reached
print("Early stopping: {} was best AUC at epoch {}".format(best_auc, best_epoch))
break
if epoch % 10 == 0:
save_model(model, resulting_folder,
"{}_fold={}_subfold={}_frac={}_n={}".format(model_name, fold, subfold, frac_string, epoch))
if do_eval:
with open(
f'{resulting_folder}/results_{model_name}_fold{fold}_subfold{subfold}_frac{frac_string}-{epoch}',
'wb') as result_file:
pickle.dump(results_dict, result_file)
else:
title_suffix = 'Final (Evaluation) EPOCH'
# check performance on test set
test_loss, final_results_dict = eval_epoch_subjects(model, criterion, test_loader, DEVICE, values, 'Test_',
plot_tpr_fpr=resulting_folder + f"roc_{model_name}_{fold}.png",
scaler=scaler, title_suffix=title_suffix)
y_true, y_pred, y_probs = final_results_dict["y_true"], final_results_dict["y_pred"], final_results_dict["y_prob"]
sens, spec = eva.compute_sensitivity_specificity(y_true, y_pred, y_scores=y_probs, prefix='Test')
roc_results = eva.compute_roc_curve(y_true, y_probs, plot=True, title_suffix='Evaluation ')
#roc = roc_results['roc'], fpr = roc_results['fpr'], tpr = roc_results['tpr'], thresholds = roc_results['thresholds'], best_index = roc_results['best_index_val']
roc_auc = roc_results['roc_auc_val'] if "roc_auc_val" in roc_results.keys() else math.nan
wandb.log({"sensitivity/Test": sens,
"specificity/Test": spec,
"AUC_ROC/Test": float(roc_auc)})
print(f"Finished with test AUC (Youden) of {roc_auc}.")
# save model parameters and predictions
save_model(model, resulting_folder,
"{}_fold={}_subfold={}_frac={}_FINAL".format(model_name, fold, subfold, frac_string))
with open(f'{resulting_folder}/results_{model_name}_fold={fold}_frac={frac_string}', 'wb') as result_file:
pickle.dump(final_results_dict, result_file)