SwinUNETR / Segmentation /train.py
deboraJ23's picture
uploaded files from https://github.com/smriti-joshi/bcnaim-odelia-challenge (except Readme, Licence and .gitignore)
361b108 verified
Raw
History Blame Contribute Delete
6.11 kB
import numpy as np
import pandas as pd
from torch.utils.data import DataLoader
import torch
import torch.optim as optim
from tqdm import tqdm
import yaml
import wandb
import math
from torch.optim.lr_scheduler import LambdaLR
from monai.losses import DiceCELoss
from monai.metrics import DiceMetric
from monai.transforms import Activations, AsDiscrete
from monai.networks.nets import SwinUNETR # <-- Ensure you have the correct import for your modelS
from dataloading.dataloader2D import NiftiSegmentationDataset # <-- You must define this properly!
from monai.transforms import (
Activations,
AsDiscrete,
Compose)
import cv2
# -------------------------
# Warmup Cosine Scheduler
# -------------------------
def warmup_cosine_lr_scheduler(optimizer, warmup_epochs, total_epochs):
def lr_lambda(current_epoch):
if current_epoch < warmup_epochs:
return float(current_epoch) / float(max(1, warmup_epochs))
else:
return 0.5 * (1. + math.cos(math.pi * (current_epoch - warmup_epochs) / (total_epochs - warmup_epochs)))
return LambdaLR(optimizer, lr_lambda)
# -------------------------
# Load Config
# -------------------------
with open("/workspace/Segmentation/config2d.yaml", "r") as f:
config = yaml.safe_load(f)
wandb.init(project=config['project_name'], config=config, name=config['run_name'], notes=config['notes'])
cfg = wandb.config
device = torch.device(cfg.device if torch.cuda.is_available() else "cpu")
# -------------------------
# Model Setup
# -------------------------
in_channels = len(cfg.channel_keys) if isinstance(cfg.channel_keys, list) else 1
model =SwinUNETR(img_size=(256, 256), in_channels=in_channels, out_channels=cfg.num_classes, spatial_dims=2).to(device)
optimizer = optim.Adam(model.parameters(), lr=cfg.learning_rate, weight_decay=1e-5)
scheduler = warmup_cosine_lr_scheduler(optimizer, cfg.warmup_epochs, cfg.epochs)
criterion = DiceCELoss(sigmoid=True, to_onehot_y=True)
post_pred = Compose([Activations(sigmoid=True), AsDiscrete(threshold=0.5)])
dice_metric = DiceMetric(include_background=False, reduction="mean", get_not_nans=False)
# -------------------------
# Dataloaders
# -------------------------
train_dataset = NiftiSegmentationDataset(cfg.csv_file_train, channel_keys=cfg.channel_keys)
train_loader = DataLoader(train_dataset, batch_size=cfg.batch_size, shuffle=True)
val_dataset = NiftiSegmentationDataset(cfg.csv_file_val, channel_keys=cfg.channel_keys, augment=False)
val_loader = DataLoader(val_dataset, batch_size=cfg.batch_size, shuffle=False)
best_val_loss = float('inf')
# -------------------------
# Training Loop
# -------------------------
for epoch in range(cfg.epochs):
model.train()
total_loss = 0.0
for batch in tqdm(train_loader, desc=f"Epoch {epoch+1}/{cfg.epochs}"):
x, y = batch
x, y = x.to(device), y.to(device)
optimizer.zero_grad()
logits, _ = model(x)
loss = criterion(logits, y)
loss.backward()
optimizer.step()
total_loss += loss.item() * x.size(0)
avg_train_loss = total_loss / len(train_loader.dataset)
# -------------------------
# Validation
# -------------------------
model.eval()
val_loss = 0.0
dice_metric.reset()
with torch.no_grad():
for batch in tqdm(val_loader, desc="Validation"):
x_val, y_val = batch
x_val, y_val = x_val.to(device), y_val.to(device)
val_logits, _ = model(x_val)
v_loss = criterion(val_logits, y_val)
val_loss += v_loss.item() * x_val.size(0)
val_probs = post_pred(val_logits)
dice_metric(y_pred=val_probs, y=y_val)
avg_val_loss = val_loss / len(val_loader.dataset)
avg_dice = dice_metric.aggregate().item()
dice_metric.reset()
# -------------------------
# Logging & Saving
# -------------------------
print(f"Epoch {epoch+1}: Train Loss={avg_train_loss:.4f} | Val Loss={avg_val_loss:.4f} | Dice={avg_dice:.4f}")
if epoch % 30 == 0:
num_examples = min(5, x_val.shape[0])
for i in range(num_examples):
img_np = x_val[i][0].detach().cpu().numpy()
gt_np = y_val[i].detach().cpu().numpy().astype(np.uint8)
pred_np = torch.argmax(val_probs[i], dim=0).detach().cpu().numpy().astype(np.uint8)
# Squeeze image channel if needed
if img_np.shape[0] == 1:
img_np = img_np[0]
# Squeeze and cast to uint8
gt_np = np.squeeze(gt_np).astype(np.uint8)
pred_np = np.squeeze(pred_np).astype(np.uint8)
# Normalize and convert to RGB
img_norm = (img_np - img_np.min()) / (img_np.max() - img_np.min() + 1e-8)
img_uint8 = (img_norm * 255).astype(np.uint8)
img_rgb = cv2.cvtColor(img_uint8, cv2.COLOR_GRAY2BGR)
# Clone image for drawing
contour_img = img_rgb.copy()
# Draw GT contours in green
contours_gt, _ = cv2.findContours(gt_np, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
cv2.drawContours(contour_img, contours_gt, -1, (0, 255, 0), 2)
# Draw predicted contours in red
contours_pred, _ = cv2.findContours(pred_np, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
cv2.drawContours(contour_img, contours_pred, -1, (0, 0, 255), 2)
# Log to wandb
wandb.log({
f"val/contours_{i}": wandb.Image(contour_img, caption=f"Epoch {epoch+1} - GT (green), Pred (red)")
})
wandb.log({
"epoch": epoch + 1,
"train_loss": avg_train_loss,
"val_loss": avg_val_loss,
"val_dice": avg_dice,
"lr": scheduler.get_last_lr()[0]
})
scheduler.step()
torch.save(model.state_dict(), '/workspace/Segmentation/checkpoints/latest_model.pth')
if avg_val_loss < best_val_loss:
best_val_loss = avg_val_loss
torch.save(model.state_dict(), '/workspace/Segmentation/checkpoints/best_model.pth')
print("✅ Saved best model.")