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}") # Load label file, only use odelia for now 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) # Get tio datasets and loaders 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 ) # Loss function and optimizer. 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 = [], [] # Same some results for later debugging 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}")): # Forward pass 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() # Gradient accumulation step 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 # Handle remaining gradients 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"]) # Validation 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() # Update training progress 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") # Clean up gc.collect() torch.cuda.empty_cache() return if __name__ == "__main__": main()