| import argparse |
| import json |
| import pandas as pd |
| import numpy as np |
| from tqdm import tqdm |
| from pathlib import Path |
| import torch |
| import torch.nn as nn |
| import matplotlib |
| matplotlib.use("Agg") |
| import random |
| import warnings |
| warnings.filterwarnings("ignore", category=UserWarning, module="torchio.data.image") |
| import os |
| import gc |
|
|
| from torch.optim.lr_scheduler import CosineAnnealingLR |
| from torch.amp import GradScaler, autocast |
| os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" |
| torch.multiprocessing.set_sharing_strategy("file_system") |
|
|
| from model.resnet import Resnet |
| from utils.dataloader import prepare_loaders, prepare_batch_single_scan |
| from utils.metrics import compute_metrics_mc |
| from utils.plots import plot_training_progress_classification, plot_pred_summary_mc |
| from utils.utils import compute_class_weights |
|
|
|
|
| def collect_arguments(): |
| """ |
| Collects arguments from the command line. |
| """ |
| |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--df", type=Path, required=True) |
| parser.add_argument("--savedir", type=Path, required=True) |
| parser.add_argument("--model_name", type=str, default="resnet18", required=True) |
| parser.add_argument("--optimizer", type=str, choices=["adamw", "sgd"], default="adamw", required=True) |
| parser.add_argument("--lr", type=float, default=1e-5, required=True) |
| parser.add_argument("--augment", action="store_true", default=False) |
| parser.add_argument("--accum_steps", type=int, default=4, required=True) |
| parser.add_argument("--batch_size", type=int, default=4, required=True) |
| parser.add_argument("--cnn_wd", type=float, default=1e-3, required=True) |
| parser.add_argument("--debug_subset", action="store_true", default=False) |
| parser.add_argument("--balanced_sampling", action="store_true", default=False) |
| parser.add_argument("--val_fold", type=int, default=4, required=True) |
| args = parser.parse_args() |
| |
| assert args.df.exists(), "Dataset file does not exist" |
| |
| return args.df, args.savedir, args.model_name, args.optimizer, args.lr, args.augment, args.accum_steps, args.batch_size, args.cnn_wd, args.debug_subset, args.balanced_sampling, args.val_fold |
|
|
|
|
| def main(): |
| """ |
| Baseline method that uses a simple resnet with a single scan as input. |
| """ |
| |
| df, savedir, model_name, optimizer_name, lr, do_augmentation, accum_steps, batch_size, cnn_wd, debug_subset, balanced_sampling, val_fold = collect_arguments() |
| |
| run_id = ''.join([random.choice('0123456789abcdef') for _ in range(6)]) |
| save_dir = Path(str(savedir).replace("run_id", run_id)) |
| save_dir.mkdir(parents=True, exist_ok=True) |
| save_dir.joinpath("preds").mkdir(parents=True, exist_ok=True) |
| save_dir.joinpath("model_weights").mkdir(parents=True, exist_ok=True) |
| print(f"Saving results to {save_dir}") |
| |
| |
| df = pd.read_excel(df) |
| df = df[(df["scan_number"] == "mip") & (df["breast_label"] != "n.a.")] |
| df.to_excel(save_dir.joinpath("train_data.xlsx"), index=False) |
|
|
| |
| train_loader, val_loader, _, _ = prepare_loaders( |
| df=df, |
| do_augmentation=do_augmentation, |
| mode="training", |
| batch_size=batch_size, |
| all_scans=False, |
| max_num_scans=1, |
| finetune_label="breast_label", |
| debug_subset=debug_subset, |
| balanced_sampling=balanced_sampling, |
| num_classes=3, |
| val_fold=val_fold |
| ) |
|
|
| |
| n_epochs = 200 |
| if not balanced_sampling: |
| class_weights = compute_class_weights(df=df, label="breast_label", dataset="odelia", all_scans=False) |
| criterion = nn.CrossEntropyLoss(weight=class_weights.half()) |
| else: |
| criterion = nn.CrossEntropyLoss() |
| |
| model = Resnet(model_name=model_name, num_classes=3, norm="batch") |
| model.to("cuda") |
| |
| optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=cnn_wd) |
| scheduler = CosineAnnealingLR(optimizer, T_max=n_epochs, eta_min=lr/10) |
| |
| all_train_loss, all_val_loss = [], [] |
| all_train_acc, all_val_acc = [], [] |
| all_train_auc, all_val_auc = [], [] |
| |
| |
| training_config = { |
| "model_name": model_name, |
| "optimizer_name": optimizer_name, |
| "scheduler": "cosine", |
| "start_lr": lr, |
| "end_lr": lr/10, |
| "n_epochs": n_epochs, |
| "do_augmentation": do_augmentation, |
| "batch_size": batch_size, |
| "run_id": run_id, |
| "balanced_sampling": balanced_sampling, |
| "accum_steps": accum_steps, |
| "cnn_wd": cnn_wd, |
| "finetune_label": "breast_label", |
| "use_amp": True, |
| "mode": "training", |
| "norm": "batch", |
| "val_fold": val_fold |
| } |
| with open(save_dir.joinpath("training_config.json"), "w") as f: |
| json.dump(training_config, f, indent=4, sort_keys=True) |
| |
| scaler = GradScaler("cuda") |
| for epoch in range(n_epochs): |
| |
| model.train() |
| |
| train_epoch_loss = 0 |
| train_epoch_preds, train_epoch_probs, train_epoch_gt = [], [], [] |
| |
| for batch_idx, batch in enumerate(tqdm(train_loader, desc=f"Epoch {epoch+1}/{n_epochs}")): |
| |
| |
| x, y = prepare_batch_single_scan(batch) |
| |
| with autocast(device_type='cuda'): |
| y_pred = model(x) |
| train_loss = criterion(y_pred, y) |
| train_loss /= accum_steps |
| |
| scaler.scale(train_loss).backward() |
| |
| |
| if (batch_idx + 1) % accum_steps == 0: |
| scaler.step(optimizer) |
| scaler.update() |
| optimizer.zero_grad() |
| |
| probs = torch.softmax(y_pred, dim=1).detach().cpu().numpy() |
| preds = np.argmax(probs, axis=1).tolist() |
| gts = y.detach().cpu().numpy().tolist() |
| |
| train_epoch_probs.extend(probs.tolist()) |
| train_epoch_preds.extend(preds) |
| train_epoch_gt.extend(gts) |
| |
| train_epoch_loss += train_loss.item() * accum_steps |
| |
| |
| if len(train_loader) % accum_steps != 0: |
| scaler.step(optimizer) |
| scaler.update() |
| optimizer.zero_grad() |
| |
| all_train_loss.append(train_epoch_loss / len(train_loader.dataset)) |
| train_metrics = compute_metrics_mc(labels=train_epoch_gt, preds=train_epoch_preds, probs=train_epoch_probs, num_classes=3) |
| all_train_acc.append(train_metrics["balanced_accuracy"]) |
| all_train_auc.append(train_metrics["auc"]) |
| |
| |
| model.eval() |
| val_epoch_loss = 0 |
| val_epoch_preds, val_epoch_probs, val_epoch_gt = [], [], [] |
| |
| with torch.no_grad(): |
| for batch in tqdm(val_loader, desc=f"Epoch {epoch+1}/{n_epochs} - Val"): |
| |
| x, y = prepare_batch_single_scan(batch) |
| |
| with autocast(device_type='cuda'): |
| y_pred = model(x) |
| val_loss = criterion(y_pred, y) |
| |
| probs = torch.softmax(y_pred, dim=1).detach().cpu().numpy() |
| preds = np.argmax(probs, axis=1).tolist() |
| gts = y.detach().cpu().numpy().tolist() |
| |
| val_epoch_loss += val_loss.item() |
| val_epoch_probs.extend(probs.tolist()) |
| val_epoch_preds.extend(preds) |
| val_epoch_gt.extend(gts) |
| |
| all_val_loss.append(val_epoch_loss / len(val_loader.dataset)) |
| val_metrics = compute_metrics_mc(labels=val_epoch_gt, preds=val_epoch_preds, probs=val_epoch_probs, num_classes=3) |
| all_val_acc.append(val_metrics["balanced_accuracy"]) |
| all_val_auc.append(val_metrics["auc"]) |
| |
| scheduler.step() |
| |
| |
| save_path = save_dir.joinpath(f"progress.png") |
| plot_training_progress_classification( |
| all_train_loss=all_train_loss, |
| all_val_loss=all_val_loss, |
| all_train_acc=all_train_acc, |
| all_val_acc=all_val_acc, |
| all_train_auc=all_train_auc, |
| all_val_auc=all_val_auc, |
| save_path=save_path |
| ) |
| |
| if all_val_auc[-1] == max(all_val_auc): |
| torch.save(model.state_dict(), save_dir.joinpath("model_weights", "best_model.pth")) |
| |
| if epoch % 5 == 0: |
| torch.save(model.state_dict(), save_dir.joinpath("model_weights", f"epoch_{str(epoch).zfill(3)}_model.pth")) |
| plot_pred_summary_mc( |
| preds=train_epoch_preds, |
| probs=train_epoch_probs, |
| gts=train_epoch_gt, |
| n_classes=3, |
| save_path=save_dir.joinpath("preds", f"epoch_{str(epoch).zfill(3)}_train_preds.png") |
| ) |
| plot_pred_summary_mc( |
| preds=val_epoch_preds, |
| probs=val_epoch_probs, |
| gts=val_epoch_gt, |
| n_classes=3, |
| save_path=save_dir.joinpath("preds", f"epoch_{str(epoch).zfill(3)}_val_preds.png") |
| ) |
| |
| metric_df = pd.DataFrame({ |
| "epoch": list(range(1, epoch+2)), |
| "train_loss": all_train_loss, |
| "val_loss": all_val_loss, |
| "train_acc": all_train_acc, |
| "val_acc": all_val_acc, |
| "train_auc": all_train_auc, |
| "val_auc": all_val_auc |
| }) |
| metric_df.to_excel(save_dir.joinpath("metrics.xlsx"), index=False) |
| |
| print(f"Epoch {epoch+1}/{n_epochs}: > Train Loss: {all_train_loss[-1]:.4f} - Val Loss: {all_val_loss[-1]:.4f}") |
| print(f"Epoch {epoch+1}/{n_epochs}: > Train Acc: {all_train_acc[-1]:.4f} - Val Acc: {all_val_acc[-1]:.4f}") |
| print(f"Epoch {epoch+1}/{n_epochs}: > Train AUC: {all_train_auc[-1]:.4f} - Val AUC: {all_val_auc[-1]:.4f}\n") |
| |
| |
| gc.collect() |
| torch.cuda.empty_cache() |
|
|
| return |
|
|
|
|
| if __name__ == "__main__": |
| main() |