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
9.39 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 DiceLoss, DiceCELoss
from monai.metrics import DiceMetric
from monai.transforms import Activations, AsDiscrete
from models.swinunetr import SwinUNETRMultiTask
from dataloading.dataloader2D import NiftiSegmentationDataset
from monai.transforms import (
Activations,
AsDiscrete,
Compose
)
from dataloading.collate_function import custom_collate
from odelia_breast_mri.scripts.main_predict import evaluate
import torch.nn.functional as F
# -------------------------
# 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/ClassifierSegmenter/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
def confidence_loss(logits):
probs = F.softmax(logits, dim=1)
entropy = -torch.sum(probs * torch.log(probs + 1e-6), dim=1)
return torch.mean(entropy)
model =SwinUNETRMultiTask(img_size=(256, 256), in_channels=in_channels, out_seg_channels=2, out_cls_classes=3).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)
# # -------------------------
# # Compute Weights and Configure Loss function
# # -------------------------
segmentation_loss = DiceCELoss(sigmoid=True, to_onehot_y=True)
df = pd.read_csv(cfg.csv_file_train)
labels = df['label'].values
class_sample_counts = np.bincount(labels)
weights = 1.0 / class_sample_counts
weights = weights / weights.sum() # normalize
weights = torch.tensor(weights, dtype=torch.float32)
print('Weights:', weights)
classification_loss = torch.nn.CrossEntropyLoss(weight=weights.to(device))
# classification_loss = torch.nn.CrossEntropyLoss()
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, collate_fn=custom_collate, num_workers=cfg.num_workers)
val_dataset = NiftiSegmentationDataset(cfg.csv_file_val, channel_keys=cfg.channel_keys, augment=False)
val_loader = DataLoader(val_dataset, batch_size=cfg.batch_size, collate_fn=custom_collate, shuffle=False)
best_val_loss = float('inf')
best_val_score = float('-inf')
# -------------------------
# Training Loop
# -------------------------
for epoch in range(cfg.epochs):
model.train()
total_loss = 0.0
total_loss, correct, total = 0, 0, 0
all_preds = []
all_probs = []
all_targets = []
for batch in tqdm(train_loader, desc=f"Epoch {epoch+1}/{cfg.epochs}"):
x = batch['image'].to(device)
y = batch['cls_label'].to(device) if batch['cls_label'] is not None else None
has_label = batch['has_cls_label'] if batch['has_cls_label'] is not None else None
mask = batch['mask'].to(device) if batch['mask'] is not None else None
has_mask = batch['has_mask'].to(device) if batch['has_mask'] is not None else None
optimizer.zero_grad()
seg_preds, cls_preds, _ = model(x)
if y is not None:
valid_idx = has_label.nonzero(as_tuple=True)[0]
if len(valid_idx) > 0:
cls_loss = classification_loss(cls_preds[valid_idx], y[valid_idx])
else:
cls_loss = 0.0
else:
cls_loss = 0.0
# conf_loss = confidence_loss(logits=cls_preds)
if cfg.propagate_segmentation_loss:
# Segmentation loss (computed only where masks are valid)
if mask is not None:
valid_idx = has_mask.nonzero(as_tuple=True)[0]
if len(valid_idx) > 0:
seg_loss = segmentation_loss(seg_preds[valid_idx], mask[valid_idx])
else:
seg_loss = 0.0
else:
seg_loss = 0.0
loss = cls_loss + seg_loss
else:
# If segmentation loss is not propagated, only use classification loss
loss = cls_loss
loss.backward()
optimizer.step()
total_loss += loss.item() * x.size(0)
##### classification metric
if y is not None:
preds = torch.argmax(cls_preds, dim=1)
correct += (preds == y).sum().item()
total += y.size(0)
# --- Collect predictions ---
probs = torch.softmax(cls_preds, dim=1) # Probabilities per class
all_preds.append(preds.cpu().detach())
all_probs.append(probs.cpu().detach())
all_targets.append(y.cpu().detach())
# After loop, concatenate all
all_preds = torch.cat(all_preds)
all_probs = torch.cat(all_probs)
all_targets = torch.cat(all_targets)
train_accuracy = correct / total
train_auc, train_sensitivity, train_specificity = evaluate(all_targets, all_preds, all_probs)
avg_train_loss = total_loss / len(train_loader.dataset)
# -------------------------
# Validation
# -------------------------
model.eval()
val_loss, val_correct, val_total = 0, 0, 0
all_preds = []
all_probs = []
all_targets = []
with torch.no_grad():
for batch in tqdm(val_loader, desc="Validation"):
x_val = batch['image'].to(device)
y_cls = batch['cls_label'].to(device)
y_mask = batch['mask'].to(device) if batch['mask'] is not None else None
has_mask = batch['has_mask'].to(device) if batch['has_mask'] is not None else None
seg_preds, cls_preds, _ = model(x_val)
# Classification loss
cls_loss = classification_loss(cls_preds, y_cls)
val_loss += cls_loss.item() * x_val.size(0)
preds = torch.argmax(cls_preds, dim=1)
val_correct += (preds == y_cls).sum().item()
val_total += y_cls.size(0)
probs = torch.softmax(cls_preds, dim=1) # Probabilities per class
all_preds.append(preds.cpu().detach())
all_probs.append(probs.cpu().detach())
all_targets.append(y_cls.cpu().detach())
# After loop, concatenate all
all_preds = torch.cat(all_preds)
all_probs = torch.cat(all_probs)
all_targets = torch.cat(all_targets)
avg_val_loss = val_loss / val_total
val_accuracy = val_correct / val_total
val_auc, val_sensitivity, val_specificity = evaluate(all_targets, all_preds, all_probs)
avg_val_loss = val_loss / len(val_loader.dataset)
# -------------------------
# Logging & Visualization
# -------------------------
print(
f"Epoch {epoch+1} Summary:\n"
f" Train Loss : {avg_train_loss:.4f} | Val Loss : {avg_val_loss:.4f}\n"
f" Train Accuracy : {train_accuracy:.4f} | Val Accuracy : {val_accuracy:.4f}\n"
f" Train AUC : {train_auc:.4f} | Val AUC : {val_auc:.4f}\n"
f" Train Sensitivity: {train_sensitivity:.4f} | Val Sensitivity: {val_sensitivity:.4f}\n"
f" Train Specificity: {train_specificity:.4f} | Val Specificity: {val_specificity:.4f}"
)
wandb.log({
"epoch": epoch + 1,
"train_loss": avg_train_loss,
"train_accuracy": train_accuracy,
"train_auc": train_auc,
"train_sensitivity": train_sensitivity,
"train_specificity": train_specificity,
"val_accuracy": val_accuracy,
"val_auc": val_auc,
"val_sensitivity": val_sensitivity,
"val_specificity": val_specificity,
"val_loss": avg_val_loss,
"lr": scheduler.get_last_lr()[0]
})
scheduler.step()
mean_score = (val_auc + val_sensitivity + val_specificity)/3
torch.save(model.state_dict(), '/workspace/Classifier/checkpoints/latest_model.pth')
if avg_val_loss < best_val_loss:
best_val_loss = avg_val_loss
torch.save(model.state_dict(), '/workspace/Classifier/checkpoints/best_model.pth')
print("✅ Saved best model.")
if mean_score > best_val_score:
best_val_score = mean_score
torch.save(model.state_dict(), '/workspace/Classifier/checkpoints/best_score_model.pth')
print("✅ Saved best score model.")