Download training_code/train.py from ODELIA-AI/Pimed: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/ODELIA-AI/Pimed/resolve/main/training_code/train.py
- Command line
-
hf download hf://ODELIA-AI/Pimed/training_code/train.py
-
curl -L -o train.py https://huggingface.co/ODELIA-AI/Pimed/resolve/main/training_code/train.py
10.3 kB
| 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() |