import argparse import pickle import gc from pathlib import Path import numpy as np import pandas as pd import torch.utils.data import wandb from torch import GradScaler, nn, optim import auto_detect_breast_mri.evaluation.metrics as evaluator from auto_detect_breast_mri.config import get_config, resolve_path from auto_detect_breast_mri.models.resnets import model_names from auto_detect_breast_mri.data import loaders from auto_detect_breast_mri.evaluation.gradcam import render_heatmap, select_uka_heatmap_cases CPU = "cpu" GPU = "cuda" def train_epoch(data_loader, suffix): model.train() torch.set_grad_enabled(True) epoch_loss = 0 for batch_id, batch in enumerate(data_loader): data = batch['image']['data'].to(DEVICE) ground_truth = batch['label'] if not torch.is_tensor(ground_truth): ground_truth = torch.tensor(ground_truth) ground_truth = ground_truth.to(DEVICE) optimizer.zero_grad() with torch.autocast(device_type=DEVICE, dtype=torch.float16): pred_probs = model(data) loss = criterion(pred_probs[:, 0].float(), ground_truth.float()) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() '''auc_roc["train_"].update(pred_probs[:, 0], ground_truth) acc["train_"].update(pred_probs[:, 0], ground_truth) sens["train_"].update(pred_probs[:, 0], ground_truth) spec["train_"].update(pred_probs[:, 0], ground_truth) ''' wandb.log({ "loss/train": loss.item() }) epoch_loss += loss.item() print('Train step loss: {}'.format(loss.item())) avg_epoch_loss = epoch_loss / len(data_loader) print('Train {}: \tAverage Loss: {:.6f}'.format( suffix, avg_epoch_loss)) '''wandb_dict = {} for name, value in [("Accuracy", acc["train_"]), ("AUC_ROC", auc_roc["train_"]), ("sensitivity", sens["train_"]), ("specificity", spec["train_"])]: wandb_dict[f"{name}/train"] = value.compute() value.reset() wandb.log(wandb_dict)''' return avg_epoch_loss @torch.no_grad() def eval_epoch(data_loader, suffix: str): set_length = len(data_loader) y_pred = torch.zeros((set_length, batch_size), dtype=torch.int) # predicted labels y_true = torch.zeros((set_length, batch_size), dtype=torch.int) # true labels y_probs = torch.zeros((set_length, batch_size), dtype=torch.float) # predicted probabilities for cancer epoch_loss = 0 for batch_id, batch in enumerate(data_loader): gc.collect() optimizer.zero_grad() data = batch['image']['data'].to(DEVICE) ground_truth = batch['label'] if not torch.is_tensor(ground_truth): ground_truth = torch.tensor(ground_truth) ground_truth = ground_truth.to(DEVICE) with torch.autocast(device_type=DEVICE, dtype=torch.float16): pred_probs = model(data) prediction = pred_probs[:, 0] > 0 loss = criterion(pred_probs[:, 0].float(), ground_truth.float()) # store prediction results for ROC y_true[batch_id][0:len(ground_truth)] = ground_truth y_pred[batch_id][0:len(ground_truth)] = prediction y_probs[batch_id][0:len(ground_truth)] = pred_probs[:, 0] epoch_loss += loss.item() print('Validation step loss: {}'.format(loss.item())) avg_epoch_loss = epoch_loss / len(data_loader) print('Validation {}: \n\tAverage Loss: {:.6f}'.format( suffix, avg_epoch_loss)) return avg_epoch_loss, y_pred, y_probs, y_true def render_gradcam_heatmaps(model, model_key, dataset, output_path, device, per_label=1): """ Save GradCAM heatmaps of `model` for a deterministic selection of test samples: the first `per_label` samples per label, ordered by patient key. One PNG per case is written to /heatmaps and logged to wandb. :return: list of written png paths """ heatmap_dir = Path(output_path) heatmap_dir.mkdir(parents=True, exist_ok=True) written = [] for index, label in select_uka_heatmap_cases(dataset, per_label=per_label): subject = dataset[index] patient_key = subject['path'] volume = subject['image']['data'].unsqueeze(0) # (1, C, W, H, D), torchio's axis order png = heatmap_dir / f"{model_key}_label{label}_{patient_key}.png" # render_heatmap never raises; it logs and skips if GradCAM is unavailable render_heatmap(model, model_key, volume, str(png), device=device) if png.exists(): print(f"Wrote heatmap {png}") wandb.log({f"heatmaps/{model_key}": wandb.Image(str(png), caption=f"{patient_key} (label {label})")}) written.append(png) else: print(f"No heatmap produced for {patient_key} ({model_key}).") return written def load_pretrained_model(model, model_path, DEVICE): checkpoint = torch.load(model_path, map_location=DEVICE) # torch.device(DEVICE)) state_dict = {} if 'state_dict' in checkpoint.keys(): print("read state_dict.") for k, v in checkpoint['state_dict'].items(): if k.startswith('module.'): key = k.replace('module.', '') else: key = k # check dimension of checkpoint value: desired_shape = model.state_dict()[key].shape if desired_shape != v.shape: pretrained_layer_tensor = torch.empty(desired_shape) (l, m, n, o, p) = desired_shape (a, b, c, d, e) = v.shape # given pretrained parameters shape pretrained_layer_tensor[:a, :b, :c, :d, :e] = v # Duplicate values along the additional dimensions if a < l: for i in range(l): pretrained_layer_tensor[a + i:a + i + 1, :, :, :, :] = pretrained_layer_tensor[:a, :, :, :, :] if b < m: for j in range(m): pretrained_layer_tensor[:, b + j:b + j + 1, :, :, :] = pretrained_layer_tensor[:, :b, :, :, :] if c < n: for k in range(n): pretrained_layer_tensor[:, :, c + k:c + k + 1, :, :] = pretrained_layer_tensor[:, :, :c, :, :] if d < o: for l in range(o): pretrained_layer_tensor[:, :, :, d + l:d + l + 1, :] = pretrained_layer_tensor[:, :, :, :d, :] if e < p: for m in range(p): pretrained_layer_tensor[:, :, :, :, e + m:e + m + 1] = pretrained_layer_tensor[:, :, :, :, :e] for k, v in model.state_dict().items(): if k not in state_dict.keys(): state_dict[k] = model.state_dict()[k] model.load_state_dict(state_dict) return model else: model.load_state_dict(checkpoint) return model if __name__ == '__main__': torch.manual_seed(31) roc_range = [0.5, 0.5] pre_image_shape = (256, 256, 32) parser = argparse.ArgumentParser( prog="aiMRI", description="Train a certain ResNet to classify malignity of Breast MRIs." ) # required arguments # choices=["resnet50_full", "resnet50_abrv", "resnet18_full", "resnet18_abrv", "resnet18_sub", "resnet50_sub"] # output channel are full -> 6, abrv -> 3, sub -> 1, resnet18_d0_t2 -> 2 parser.add_argument("model_name", help="Enter the models name to be used.") # The three paths fall back to the site config when omitted (see config.example.yaml). parser.add_argument("data_path", nargs='?', default=None, help="Path to the root folder of the (to the size 32x256x256) cropped MRIs. " "Config key: data_root.") parser.add_argument("feature_path", nargs='?', default=None, help="Path to .xlsx/.csv file that contains label information for the desired criteria and for " "each data in data_path. Config key: metadata_file.") parser.add_argument("split_files_folder", nargs='?', default=None, help="Folder holding the per-fold split files. Config key: split_root.") # optional arguments parser.add_argument("-b", "--batch_size", type=int, help="Specifies batch size to be user. Default is 4.") parser.add_argument("-c", "--fold", type=int, help="Fold to be used.") parser.add_argument("-f", "--fraction", type=float, help="Fraction used for training.") parser.add_argument("-l", "--learning_rate", type=float, help="Learning rate to be used as initial value. Default is 0.0001.") parser.add_argument("-m", "--model_path", type=str, help="Path to pth-file of the pretrained model. " "Architecture must match the one defined by model_name.") parser.add_argument("-o", "--output_path", type=str, default=None, help="Folder in which results and GradCAM heatmaps are stored. Config key: output_root.") parser.add_argument("-g", "--heatmaps_per_label", type=int, default=1, help="Number of GradCAM heatmaps to render per label. 0 disables heatmaps. Default is 1.") 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") output_path = resolve_path(args.output_path, "output_root", "output folder") # only the leaf folder name, so no local path ends up in the run config split_files_folder_name = Path(split_files_folder.rstrip('/')).name # Read optional arguments if given batch_size = 4 if args.batch_size and args.batch_size > 0: batch_size = args.batch_size fold = 0 if args.fold: fold = args.fold # train.py writes an explicit fraction into every checkpoint name, full data included # (frac_string turns the internal -1.0 sentinel back into 1.0), so the suffix is always # present. Keep the float form: '_frac=1' would not match a '_frac=1.0_FINAL.pth'. fraction = float(args.fraction) if args.fraction else 1.0 model_path = None if args.model_path: if len(args.model_path) > 4 and args.model_path.endswith('.pth'): model_path = args.model_path else: raise ValueError("Invalid model path: " + args.model_path) else: raise ValueError( "Model path not specified. Please specify the path to the pretrained model to be used if " "no training shall be performed. ") DEVICE = GPU if not torch.cuda.is_available(): DEVICE = CPU learning_rate = 0.0001 weight_decay = 1e-2 if args.learning_rate and 0 < args.learning_rate <= 0.01: learning_rate = args.learning_rate # ------------ Initialize Model, Loss Function, Optimizer ------------ models = {model_name + '_abrv': model_names.get(model_name + '_abrv'), model_name + '_full': model_names.get(model_name + '_full')} wandb.init(**get_config().wandb_init_kwargs(), config={ "learning-rate": learning_rate, "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, "weight decay": weight_decay, "Machine": "HPC", "Mixed Precision": "True", "Fold": fold, "splits": split_files_folder_name, "fraction":fraction, }, name="{}_comparison_split={}".format(model_name, fold)) scaler = GradScaler() model_predictions = {model_name + '_abrv': [], model_name + '_full': []} total_ground_truth = None features = pd.read_csv(feature_path) if feature_path.endswith(".csv") else pd.read_excel(feature_path) for model_key, model in models.items(): protocol = model_key.split("_")[-1] protocol = "abbreviated" if protocol == "abrv" else protocol print("Get dataloader for " + protocol) test_loader = loaders.get_subjects_dataloader(path_base, features, pre_image_shape, transform=None, protocol=protocol, data_selection_file_sceleton=str(Path(split_files_folder)/f"fold{fold}/stratified_test_set"), batch_size=batch_size, fraction=1, # here we do not use fraction as for the test set we always use 100% split fold=fold, subfold=None, segmentation_task=False) wandb.log({"number of test samples abrv": len(test_loader.dataset)}) model.to(DEVICE) criterion = nn.BCEWithLogitsLoss() optimizer = optim.AdamW(model.parameters(), lr=0.0001, weight_decay=weight_decay) # load pretrained model: print(model_path) frac_suffix = f"_frac={fraction}" current_model_path = model_path.format(model_key=model_key, fold=fold, frac_suffix=frac_suffix) print("load " + current_model_path) model = load_pretrained_model(model, current_model_path, DEVICE) model.eval() # check performance on test set test_loss, y_pred, y_probs, y_true = eval_epoch(test_loader, model_key) model_predictions[model_key] = y_probs if total_ground_truth is None: total_ground_truth = y_true else: label_match = sum(total_ground_truth) == sum(y_true) print(f"Model {model_name} has the same amount of true label in the dataset: {label_match}.") sens, spec = evaluator.compute_sensitivity_specificity(y_true, y_pred, y_scores=y_probs, prefix='Test') roc_values = evaluator.compute_roc_curve(y_true, y_probs, plot=True, title_suffix='Test ' + model_key) fpr = roc_values.get("fpr", None) tpr = roc_values.get("tpr", None) roc_auc = roc_values.get("roc_auc_val", None) wandb.log( {"average_loss/Test": test_loss, "sensitivity/Test": sens, "specificity/Test": spec, "AUC_ROC/Test": roc_auc}) print(f"Finished with test AUC(Youden) {roc_auc}.") results_dict = {"y_pred": y_pred, "y_probs": y_probs, "y_true": y_true} with open(f'{output_path}/results_{model_key}_fold={fold}_frac={fraction}', 'wb') as result_file: #or split_test_{model_key}_{fold}_f{fraction}.txt results_{model_key}_comparison_fold{fold}', 'wb') as result_file: pickle.dump(results_dict, result_file) # deterministic heatmaps: first sample(s) per label of this protocol's test set if args.heatmaps_per_label > 0: render_gradcam_heatmaps(model, model_key, test_loader.dataset, Path(output_path) / "heatmaps" /f"fold{fold}-fract{fraction}", DEVICE, per_label=args.heatmaps_per_label) abrv_predictions = model_predictions.get(model_name + '_abrv').detach().numpy() full_predictions = model_predictions.get(model_name + '_full').detach().numpy() p_value = evaluator.delong_roc_test(np.array(total_ground_truth), abrv_predictions, full_predictions) print(f"Models {model_predictions.keys()} have the p value {p_value}") wandb.log({"p_value": p_value})