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.")