| 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 |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| 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") |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| 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() |
| weights = torch.tensor(weights, dtype=torch.float32) |
| print('Weights:', weights) |
|
|
| classification_loss = torch.nn.CrossEntropyLoss(weight=weights.to(device)) |
| |
|
|
|
|
| post_pred = Compose([Activations(sigmoid=True), AsDiscrete(threshold=0.5)]) |
| dice_metric = DiceMetric(include_background=False, reduction="mean", get_not_nans=False) |
|
|
| |
| |
| |
| 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') |
|
|
| |
| |
| |
| 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 |
|
|
| |
|
|
| if cfg.propagate_segmentation_loss: |
| |
| 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: |
| |
| loss = cls_loss |
| loss.backward() |
| optimizer.step() |
|
|
| total_loss += loss.item() * x.size(0) |
|
|
| |
| if y is not None: |
| preds = torch.argmax(cls_preds, dim=1) |
| correct += (preds == y).sum().item() |
| total += y.size(0) |
| |
| probs = torch.softmax(cls_preds, dim=1) |
| all_preds.append(preds.cpu().detach()) |
| all_probs.append(probs.cpu().detach()) |
| all_targets.append(y.cpu().detach()) |
|
|
| |
| 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) |
|
|
| |
| |
| |
| 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) |
|
|
| |
| 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) |
|
|
| all_preds.append(preds.cpu().detach()) |
| all_probs.append(probs.cpu().detach()) |
| all_targets.append(y_cls.cpu().detach()) |
|
|
|
|
| |
| 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) |
|
|
| |
| |
| |
| 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.") |