Pimed / training_code /train.py
deboraJ23's picture
upload training_code
64fd08f verified
Raw
History Blame Contribute Delete
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()