Download Classification/train.py from ODELIA-AI/SwinUNETR: direct link, hf CLI and curl.
- Browser
- Download file 9.39 kB
-
https://huggingface.co/ODELIA-AI/SwinUNETR/resolve/main/Classification/train.py
- Command line
-
hf download hf://ODELIA-AI/SwinUNETR/Classification/train.py
-
curl -L -o train.py https://huggingface.co/ODELIA-AI/SwinUNETR/resolve/main/Classification/train.py
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.") |