Download scripts/validation/compare_protocols.py from deboraJ23/AI_MRI: direct link, hf CLI and curl.
- Browser
- Download file 16.3 kB
-
https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/validation/compare_protocols.py
- Command line
-
hf download hf://deboraJ23/AI_MRI/scripts/validation/compare_protocols.py
-
curl -L -o compare_protocols.py https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/validation/compare_protocols.py
16.3 kB
| 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 | |
| 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 <output_path>/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}) | |