Image Classification
timm
English
medical-imaging
knee-mri
acl-tear-detection
deep-learning
convnext
self-attention
masked-slice-modeling
radiology
orthopedics
Eval Results (legacy)
Instructions to use shareefch1413/ACL-LKNet with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- timm
How to use shareefch1413/ACL-LKNet with timm:
import timm model = timm.create_model("hf-hub:shareefch1413/ACL-LKNet", pretrained=True) - Notebooks
- Google Colab
- Kaggle
Download src/train.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 27.5 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/train.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/src/train.py
-
curl -L -o train.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/train.py
27.5 kB
| """ | |
| Training loops for ACL-LKNet. | |
| Two phases: | |
| Phase 1: SSL Pretraining (Masked Slice Modeling) | |
| - Train backbone to reconstruct masked slice features | |
| - Monitor pretext loss, stop when plateaus | |
| Phase 2: Supervised Fine-tuning | |
| - Load SSL-pretrained backbone | |
| - Train full model with differential LR | |
| - Weighted BCE loss for class imbalance | |
| - Early stopping on validation AUROC | |
| Both phases support: | |
| - Mixed precision (FP16) for T4 memory | |
| - Gradient accumulation for effective batch size | |
| - Gradient clipping for stability | |
| - EMA model for better generalization | |
| - Full checkpoint save/load for Colab session recovery | |
| """ | |
| import os | |
| import copy | |
| import time | |
| import math | |
| import logging | |
| from typing import Optional, Dict, Tuple, Any, List | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| from torch.utils.data import DataLoader | |
| from tqdm import tqdm | |
| from .config import Config | |
| from .models.acl_lknet import ACLLKNet, create_model_from_config | |
| from .models.msm import MaskedSliceModeling | |
| from .dataset import create_dataloaders | |
| from .utils import ( | |
| set_seed, EMAModel, save_checkpoint, load_checkpoint, | |
| find_latest_checkpoint, setup_logging, format_metrics, | |
| get_gpu_memory_info, clear_gpu_memory, | |
| ) | |
| from .evaluate import compute_metrics | |
| # ── Loss Functions ────────────────────────────────────────────────── | |
| def create_loss_fn(config: Config, device: torch.device = None) -> nn.Module: | |
| """Create weighted BCE loss with label smoothing. | |
| Args: | |
| config: Config with pos_weight setting | |
| device: Target device for pos_weight tensor (avoids CPU/GPU mismatch) | |
| """ | |
| pos_weight = torch.tensor([config.pos_weight], device=device) | |
| return nn.BCEWithLogitsLoss(pos_weight=pos_weight) | |
| def apply_label_smoothing(labels: torch.Tensor, smoothing: float = 0.05) -> torch.Tensor: | |
| """Apply label smoothing: 0 → smoothing, 1 → 1-smoothing.""" | |
| return labels * (1 - smoothing) + 0.5 * smoothing | |
| def apply_mixup( | |
| batch: dict, alpha: float = 0.2 | |
| ) -> Tuple[dict, torch.Tensor, torch.Tensor, float]: | |
| """ | |
| Apply Mixup augmentation to a batch. | |
| Returns modified batch, original labels, shuffled labels, and lambda. | |
| """ | |
| if alpha <= 0: | |
| return batch, batch["label"], batch["label"], 1.0 | |
| lam = np.random.beta(alpha, alpha) | |
| lam = max(lam, 1 - lam) # Ensure lam >= 0.5 | |
| B = batch["sagittal"].shape[0] | |
| if B < 2: | |
| return batch, batch["label"], batch["label"], 1.0 | |
| indices = torch.randperm(B) | |
| mixed_batch = {} | |
| for key in ["sagittal", "coronal", "axial"]: | |
| mixed_batch[key] = lam * batch[key] + (1 - lam) * batch[key][indices] | |
| for key in ["sag_mask", "cor_mask", "axi_mask"]: | |
| mixed_batch[key] = batch[key] | |
| mixed_batch["label"] = batch["label"] | |
| mixed_batch["case_id"] = batch["case_id"] | |
| return mixed_batch, batch["label"], batch["label"][indices], lam | |
| # ── Phase 1: SSL Pretraining ────────────────────────────────────── | |
| def pretrain_msm(config: Config) -> str: | |
| """ | |
| Self-supervised pretraining via Masked Slice Modeling. | |
| Returns path to the best checkpoint. | |
| """ | |
| setup_logging(config.log_dir) | |
| set_seed(config.seed) | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| logging.info(f"MSM Pretraining | Device: {device} | Strategy: {config.mask_strategy}") | |
| logging.info(f"GPU: {get_gpu_memory_info()}") | |
| # Data | |
| train_loader, val_loader = create_dataloaders(config, ssl=True) | |
| logging.info(f"Train: {len(train_loader.dataset)} exams | Val: {len(val_loader.dataset)} exams") | |
| # Model | |
| model = create_model_from_config(config).to(device) | |
| msm = MaskedSliceModeling( | |
| feature_dim=model.feature_dim, | |
| decoder_dim=config.msm_decoder_dim, | |
| decoder_layers=config.msm_decoder_layers, | |
| decoder_heads=config.msm_decoder_heads, | |
| max_slices=config.max_slices, | |
| mask_ratio=config.mask_ratio, | |
| mask_strategy=config.mask_strategy, | |
| ).to(device) | |
| # Optimizer | |
| params = list(model.parameters()) + list(msm.parameters()) | |
| optimizer = torch.optim.AdamW(params, lr=config.ssl_lr, weight_decay=config.ssl_weight_decay) | |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2) | |
| scaler = torch.amp.GradScaler("cuda", enabled=config.use_amp) | |
| # Checkpoint recovery | |
| start_epoch = 0 | |
| best_loss = float("inf") | |
| train_history = [] | |
| val_history = [] | |
| patience_counter = 0 | |
| ckpt_path = find_latest_checkpoint(config.checkpoint_dir, phase="ssl") | |
| if ckpt_path: | |
| logging.info(f"Resuming SSL from checkpoint: {ckpt_path}") | |
| ckpt = load_checkpoint(ckpt_path, model, optimizer, scheduler, scaler) | |
| start_epoch = ckpt["epoch"] + 1 | |
| best_loss = ckpt["best_metric"] | |
| train_history = ckpt.get("train_history", []) | |
| val_history = ckpt.get("val_history", []) | |
| patience_counter = ckpt.get("patience_counter", 0) | |
| # Load MSM state if saved | |
| if "msm_state_dict" in ckpt: | |
| msm.load_state_dict(ckpt["msm_state_dict"]) | |
| best_ckpt_path = os.path.join(config.checkpoint_dir, "ssl_best.pt") | |
| # Training loop | |
| epoch = start_epoch | |
| while True: # No fixed epoch count — stop when loss plateaus | |
| model.train() | |
| msm.train() | |
| epoch_loss = 0.0 | |
| num_batches = 0 | |
| pbar = tqdm(train_loader, desc=f"SSL Epoch {epoch}", leave=False) | |
| optimizer.zero_grad() | |
| for step, batch in enumerate(pbar): | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| with torch.amp.autocast("cuda", enabled=config.use_amp): | |
| # Extract features | |
| feats = model.get_slice_features(sag, cor, axi) | |
| # MSM on each view independently | |
| total_loss = 0.0 | |
| for view_name in ["sagittal", "coronal", "axial"]: | |
| loss, _, _ = msm(feats[view_name]) | |
| total_loss = total_loss + loss | |
| total_loss = total_loss / 3.0 # Average over views | |
| # Gradient accumulation | |
| total_loss = total_loss / config.accumulation_steps | |
| scaler.scale(total_loss).backward() | |
| if (step + 1) % config.accumulation_steps == 0: | |
| scaler.unscale_(optimizer) | |
| torch.nn.utils.clip_grad_norm_(params, config.gradient_clip) | |
| scaler.step(optimizer) | |
| scaler.update() | |
| optimizer.zero_grad() | |
| epoch_loss += total_loss.item() * config.accumulation_steps | |
| num_batches += 1 | |
| pbar.set_postfix(loss=f"{total_loss.item() * config.accumulation_steps:.4f}") | |
| avg_loss = epoch_loss / max(num_batches, 1) | |
| scheduler.step() | |
| train_history.append({"epoch": epoch, "loss": avg_loss}) | |
| # Validation | |
| val_loss = _validate_msm(model, msm, val_loader, config, device) | |
| val_history.append({"epoch": epoch, "loss": val_loss}) | |
| logging.info( | |
| f"SSL Epoch {epoch} | Train Loss: {avg_loss:.4f} | Val Loss: {val_loss:.4f} | " | |
| f"LR: {optimizer.param_groups[0]['lr']:.2e}" | |
| ) | |
| # Check improvement | |
| if val_loss < best_loss: | |
| best_loss = val_loss | |
| patience_counter = 0 | |
| # Save best | |
| _save_ssl_checkpoint( | |
| best_ckpt_path, epoch, model, msm, optimizer, scheduler, | |
| scaler, best_loss, train_history, val_history, patience_counter, config, | |
| ) | |
| logging.info(f" ✓ New best SSL loss: {best_loss:.4f}") | |
| else: | |
| patience_counter += 1 | |
| logging.info(f" ✗ No improvement ({patience_counter}/{config.ssl_patience})") | |
| # Periodic save | |
| if (epoch + 1) % config.ssl_save_every == 0: | |
| periodic_path = os.path.join(config.checkpoint_dir, f"ssl_epoch{epoch}.pt") | |
| _save_ssl_checkpoint( | |
| periodic_path, epoch, model, msm, optimizer, scheduler, | |
| scaler, best_loss, train_history, val_history, patience_counter, config, | |
| ) | |
| # Early stopping | |
| if patience_counter >= config.ssl_patience: | |
| logging.info(f"SSL early stopping at epoch {epoch}") | |
| break | |
| epoch += 1 | |
| clear_gpu_memory() | |
| logging.info(f"SSL pretraining complete. Best loss: {best_loss:.4f}") | |
| return best_ckpt_path | |
| def _validate_msm(model, msm, val_loader, config, device) -> float: | |
| """Run MSM validation pass.""" | |
| model.eval() | |
| msm.eval() | |
| total_loss = 0.0 | |
| count = 0 | |
| with torch.no_grad(): | |
| for batch in val_loader: | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| with torch.amp.autocast("cuda", enabled=config.use_amp): | |
| feats = model.get_slice_features(sag, cor, axi) | |
| loss = 0.0 | |
| for view_name in ["sagittal", "coronal", "axial"]: | |
| l, _, _ = msm(feats[view_name]) | |
| loss = loss + l.item() | |
| loss /= 3.0 | |
| total_loss += loss | |
| count += 1 | |
| return total_loss / max(count, 1) | |
| def _save_ssl_checkpoint(path, epoch, model, msm, optimizer, scheduler, scaler, | |
| best_loss, train_history, val_history, patience_counter, config): | |
| """Save SSL checkpoint including MSM state.""" | |
| os.makedirs(os.path.dirname(path), exist_ok=True) | |
| from .utils import get_rng_states | |
| checkpoint = { | |
| "epoch": epoch, | |
| "phase": "ssl", | |
| "model_state_dict": model.state_dict(), | |
| "msm_state_dict": msm.state_dict(), | |
| "optimizer_state_dict": optimizer.state_dict(), | |
| "scheduler_state_dict": scheduler.state_dict(), | |
| "scaler_state_dict": scaler.state_dict() if scaler else None, | |
| "best_metric": best_loss, | |
| "best_epoch": epoch, | |
| "train_history": train_history, | |
| "val_history": val_history, | |
| "patience_counter": patience_counter, | |
| "rng_states": get_rng_states(), | |
| "config": config.to_dict(), | |
| } | |
| torch.save(checkpoint, path) | |
| def train_supervised( | |
| config: Config, | |
| ssl_checkpoint: Optional[str] = None, | |
| train_cases: Optional[list] = None, | |
| val_cases: Optional[list] = None, | |
| train_split: str = "train", | |
| val_split: str = "valid", | |
| ) -> str: | |
| """ | |
| Supervised fine-tuning for ACL tear detection. | |
| Args: | |
| config: Training config | |
| ssl_checkpoint: Path to SSL pretrained checkpoint (optional) | |
| train_cases: Optional explicit list of training case IDs (for CV) | |
| val_cases: Optional explicit list of validation case IDs (for CV) | |
| train_split: Dataset split directory name for training | |
| val_split: Dataset split directory name for validation | |
| Returns: | |
| Path to the best checkpoint (by val AUROC) | |
| """ | |
| setup_logging(config.log_dir) | |
| set_seed(config.seed) | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| logging.info(f"Supervised Training | Device: {device}") | |
| logging.info(f"GPU: {get_gpu_memory_info()}") | |
| # Data | |
| train_loader, val_loader = create_dataloaders( | |
| config, ssl=False, train_cases=train_cases, val_cases=val_cases, | |
| train_split=train_split, val_split=val_split, | |
| ) | |
| logging.info(f"Train: {len(train_loader.dataset)} exams | Val: {len(val_loader.dataset)} exams") | |
| # Model | |
| model = create_model_from_config(config).to(device) | |
| # Load SSL pretrained weights | |
| if ssl_checkpoint and os.path.exists(ssl_checkpoint): | |
| ssl_ckpt = torch.load(ssl_checkpoint, map_location=device, weights_only=False) | |
| model.load_state_dict(ssl_ckpt["model_state_dict"], strict=False) | |
| logging.info(f"Loaded SSL checkpoint: {ssl_checkpoint}") | |
| # Differential LR: lower for backbone, higher for new layers | |
| backbone_params = list(model.backbone.parameters()) | |
| new_params = [p for n, p in model.named_parameters() | |
| if not n.startswith("backbone")] | |
| optimizer = torch.optim.AdamW([ | |
| {"params": backbone_params, "lr": config.backbone_lr}, | |
| {"params": new_params, "lr": config.lr}, | |
| ], weight_decay=config.weight_decay) | |
| # Cosine Annealing scheduler (resume-safe — unlike OneCycleLR, does not | |
| # crash when total_steps is exceeded after checkpoint restoration) | |
| max_epochs = getattr(config, "num_epochs", 40) | |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( | |
| optimizer, | |
| T_max=max(max_epochs - config.warmup_epochs, 1), | |
| eta_min=1e-7, | |
| ) | |
| scaler = torch.amp.GradScaler("cuda", enabled=config.use_amp) | |
| criterion = create_loss_fn(config, device=device) | |
| ema = EMAModel(model, decay=config.ema_decay) | |
| # Checkpoint recovery | |
| start_epoch = 0 | |
| best_auc = 0.0 | |
| best_epoch = 0 | |
| train_history = [] | |
| val_history = [] | |
| patience_counter = 0 | |
| ckpt_path = find_latest_checkpoint(config.checkpoint_dir, phase="finetune") | |
| if ckpt_path: | |
| logging.info(f"Resuming supervised training from: {ckpt_path}") | |
| ckpt = load_checkpoint(ckpt_path, model, optimizer, scheduler, scaler, ema) | |
| start_epoch = ckpt["epoch"] + 1 | |
| best_auc = ckpt["best_metric"] | |
| best_epoch = ckpt.get("best_epoch", 0) | |
| train_history = ckpt.get("train_history", []) | |
| val_history = ckpt.get("val_history", []) | |
| patience_counter = ckpt.get("patience_counter", 0) | |
| best_ckpt_path = os.path.join(config.checkpoint_dir, "finetune_best.pt") | |
| # Training loop | |
| epoch = start_epoch | |
| while True: # Monitor-based stopping | |
| model.train() | |
| epoch_loss = 0.0 | |
| all_preds = [] | |
| all_labels = [] | |
| num_batches = 0 | |
| pbar = tqdm(train_loader, desc=f"Epoch {epoch}", leave=False) | |
| optimizer.zero_grad() | |
| for step, batch in enumerate(pbar): | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| sag_m = batch["sag_mask"].to(device) | |
| cor_m = batch["cor_mask"].to(device) | |
| axi_m = batch["axi_mask"].to(device) | |
| labels = batch["label"].to(device) | |
| # Mixup | |
| if config.mixup_alpha > 0 and epoch >= config.warmup_epochs: | |
| mixed, labels_a, labels_b, lam = apply_mixup(batch, config.mixup_alpha) | |
| sag = mixed["sagittal"].to(device) | |
| cor = mixed["coronal"].to(device) | |
| axi = mixed["axial"].to(device) | |
| labels_a = labels_a.to(device) | |
| labels_b = labels_b.to(device) | |
| else: | |
| labels_a = labels_b = labels | |
| lam = 1.0 | |
| with torch.amp.autocast("cuda", enabled=config.use_amp): | |
| output = model(sag, cor, axi, sag_m, cor_m, axi_m) | |
| logits = output["logits"].squeeze(-1) | |
| # Label smoothing | |
| smooth_a = apply_label_smoothing(labels_a, config.label_smoothing) | |
| smooth_b = apply_label_smoothing(labels_b, config.label_smoothing) | |
| # Mixup loss | |
| loss = lam * criterion(logits, smooth_a) + (1 - lam) * criterion(logits, smooth_b) | |
| loss = loss / config.accumulation_steps | |
| scaler.scale(loss).backward() | |
| if (step + 1) % config.accumulation_steps == 0: | |
| scaler.unscale_(optimizer) | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), config.gradient_clip) | |
| scaler.step(optimizer) | |
| scaler.update() | |
| optimizer.zero_grad() | |
| # Update EMA | |
| ema.update(model) | |
| epoch_loss += loss.item() * config.accumulation_steps | |
| all_preds.extend(output["probs"].squeeze(-1).detach().cpu().numpy().tolist()) | |
| all_labels.extend(labels.cpu().numpy().tolist()) | |
| num_batches += 1 | |
| pbar.set_postfix(loss=f"{loss.item() * config.accumulation_steps:.4f}") | |
| # ENH-2: Flush any remaining accumulated gradients from the tail batch | |
| # (when dataset size is not divisible by accumulation_steps) | |
| if (step + 1) % config.accumulation_steps != 0: | |
| scaler.unscale_(optimizer) | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), config.gradient_clip) | |
| scaler.step(optimizer) | |
| scaler.update() | |
| optimizer.zero_grad() | |
| ema.update(model) | |
| # Step the cosine annealing scheduler once per epoch (after warmup) | |
| if epoch >= config.warmup_epochs: | |
| scheduler.step() | |
| else: | |
| # Linear warmup: scale LR from 0 to target over warmup_epochs | |
| warmup_factor = (epoch + 1) / config.warmup_epochs | |
| for pg_idx, pg in enumerate(optimizer.param_groups): | |
| base_lr = config.backbone_lr if pg_idx == 0 else config.lr | |
| pg['lr'] = base_lr * warmup_factor | |
| avg_loss = epoch_loss / max(num_batches, 1) | |
| # Train metrics | |
| train_metrics = compute_metrics( | |
| np.array(all_labels), np.array(all_preds), prefix="train" | |
| ) | |
| train_metrics["train_loss"] = avg_loss | |
| train_history.append({"epoch": epoch, **train_metrics}) | |
| # Validation (check online model and EMA shadow model, use best) | |
| val_online = validate_supervised(model, val_loader, criterion, config, device) | |
| val_ema = validate_supervised(ema.eval_model(), val_loader, criterion, config, device) | |
| val_metrics = val_ema if val_ema["val_auroc"] >= val_online["val_auroc"] else val_online | |
| val_history.append({"epoch": epoch, **val_metrics}) | |
| logging.info( | |
| f"Ep{epoch} | Loss:{avg_loss:.4f} | " | |
| f"Online_AUROC:{val_online['val_auroc']:.4f} | EMA_AUROC:{val_ema['val_auroc']:.4f} | " | |
| f"Best_AUROC:{val_metrics['val_auroc']:.4f} | Acc:{val_metrics['val_accuracy']:.4f} | " | |
| f"Sens:{val_metrics.get('val_sensitivity', 0):.4f} | " | |
| f"LR:{optimizer.param_groups[0]['lr']:.2e}" | |
| ) | |
| # Check improvement | |
| current_auc = val_metrics["val_auroc"] | |
| if current_auc > best_auc: | |
| best_auc = current_auc | |
| best_epoch = epoch | |
| patience_counter = 0 | |
| save_checkpoint( | |
| best_ckpt_path, epoch, "finetune", model, optimizer, scheduler, | |
| scaler, ema, best_auc, best_epoch, train_history, val_history, | |
| patience_counter, config, | |
| ) | |
| logging.info(f" ✓ New best AUROC: {best_auc:.4f}") | |
| else: | |
| patience_counter += 1 | |
| logging.info(f" ✗ No improvement ({patience_counter}/{config.patience})") | |
| # Periodic save | |
| if (epoch + 1) % config.save_every == 0: | |
| periodic_path = os.path.join(config.checkpoint_dir, f"finetune_epoch{epoch}.pt") | |
| save_checkpoint( | |
| periodic_path, epoch, "finetune", model, optimizer, scheduler, | |
| scaler, ema, best_auc, best_epoch, train_history, val_history, | |
| patience_counter, config, | |
| ) | |
| # Early stopping | |
| if patience_counter >= config.patience: | |
| logging.info(f"Early stopping at epoch {epoch}. Best AUROC: {best_auc:.4f} at epoch {best_epoch}") | |
| break | |
| epoch += 1 | |
| clear_gpu_memory() | |
| logging.info(f"Training complete. Best AUROC: {best_auc:.4f} at epoch {best_epoch}") | |
| return best_ckpt_path | |
| def validate_supervised( | |
| model: nn.Module, | |
| val_loader: DataLoader, | |
| criterion: nn.Module, | |
| config: Config, | |
| device: torch.device, | |
| ) -> Dict[str, float]: | |
| """Run validation and compute metrics.""" | |
| model.eval() | |
| all_preds = [] | |
| all_labels = [] | |
| total_loss = 0.0 | |
| count = 0 | |
| with torch.no_grad(): | |
| for batch in val_loader: | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| sag_m = batch["sag_mask"].to(device) | |
| cor_m = batch["cor_mask"].to(device) | |
| axi_m = batch["axi_mask"].to(device) | |
| labels = batch["label"].to(device) | |
| with torch.amp.autocast("cuda", enabled=config.use_amp): | |
| output = model(sag, cor, axi, sag_m, cor_m, axi_m) | |
| logits = output["logits"].squeeze(-1) | |
| loss = criterion(logits, labels) | |
| total_loss += loss.item() | |
| all_preds.extend(output["probs"].squeeze(-1).cpu().numpy().tolist()) | |
| all_labels.extend(labels.cpu().numpy().tolist()) | |
| count += 1 | |
| metrics = compute_metrics(np.array(all_labels), np.array(all_preds), prefix="val") | |
| metrics["val_loss"] = total_loss / max(count, 1) | |
| return metrics | |
| # ── 5-Fold Cross-Validation Protocol ─────────────────────────────── | |
| def train_5fold_cross_validation( | |
| config: Config, | |
| ssl_checkpoint: Optional[str] = None, | |
| task: str = "acl", | |
| ) -> Dict[str, Any]: | |
| """ | |
| Execute patient-stratified 5-Fold Cross-Validation for ACL-LKNet. | |
| Partitions the dataset into 5 balanced folds, trains a separate model on each fold, | |
| evaluates out-of-fold predictions, and reports Mean ± Std across folds for all | |
| academic metrics (AUROC, Accuracy, Sensitivity, Specificity, F1, MCC). | |
| Returns: | |
| Dict with fold-by-fold results, aggregated Mean ± Std, and out-of-fold metrics. | |
| """ | |
| from .dataset import get_stratified_folds | |
| import json | |
| setup_logging(config.log_dir) | |
| logging.info(f"=== Starting Stratified {config.n_splits}-Fold Cross-Validation ===") | |
| folds = get_stratified_folds( | |
| config.data_dir, split="train", n_splits=config.n_splits, seed=config.seed, task=task | |
| ) | |
| base_exp_name = config.experiment_name | |
| fold_summaries = [] | |
| all_oof_labels = [] | |
| all_oof_preds = [] | |
| for fold_info in folds: | |
| fold_idx = fold_info["fold"] | |
| logging.info(f"\n--- Running Fold {fold_idx + 1} / {config.n_splits} ---") | |
| logging.info(f"Train Cases: {len(fold_info['train_cases'])} | Val Cases: {len(fold_info['val_cases'])}") | |
| # Create fold-specific configuration | |
| fold_config = copy.deepcopy(config) | |
| fold_config.experiment_name = f"{base_exp_name}_fold{fold_idx}" | |
| fold_config.__post_init__() | |
| os.makedirs(fold_config.checkpoint_dir, exist_ok=True) | |
| os.makedirs(fold_config.log_dir, exist_ok=True) | |
| # Train supervised on this fold | |
| best_ckpt = train_supervised( | |
| fold_config, | |
| ssl_checkpoint=ssl_checkpoint, | |
| train_cases=fold_info["train_cases"], | |
| val_cases=fold_info["val_cases"], | |
| train_split="train", | |
| val_split="train", | |
| ) | |
| # Evaluate best model on fold validation set | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| model = create_model_from_config(fold_config).to(device) | |
| load_checkpoint(best_ckpt, model) | |
| model.eval() | |
| _, fold_val_loader = create_dataloaders( | |
| fold_config, ssl=False, | |
| val_cases=fold_info["val_cases"], | |
| val_split="train", | |
| ) | |
| fold_preds = [] | |
| fold_labels = [] | |
| with torch.no_grad(): | |
| for batch in fold_val_loader: | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| sag_m = batch["sag_mask"].to(device) | |
| cor_m = batch["cor_mask"].to(device) | |
| axi_m = batch["axi_mask"].to(device) | |
| with torch.amp.autocast("cuda", enabled=fold_config.use_amp and torch.cuda.is_available()): | |
| out = model(sag, cor, axi, sag_m, cor_m, axi_m) | |
| fold_preds.extend(out["probs"].squeeze(-1).cpu().numpy().tolist()) | |
| fold_labels.extend(batch["label"].numpy().tolist()) | |
| fold_m = compute_metrics(np.array(fold_labels), np.array(fold_preds)) | |
| fold_summaries.append({ | |
| "fold": fold_idx + 1, | |
| "best_checkpoint": best_ckpt, | |
| "accuracy": fold_m["accuracy"], | |
| "balanced_accuracy": fold_m["balanced_accuracy"], | |
| "auroc": fold_m["auroc"], | |
| "auprc": fold_m["auprc"], | |
| "sensitivity": fold_m["sensitivity"], | |
| "specificity": fold_m["specificity"], | |
| "f1": fold_m["f1"], | |
| "mcc": fold_m["mcc"], | |
| }) | |
| all_oof_labels.extend(fold_labels) | |
| all_oof_preds.extend(fold_preds) | |
| # Compute Mean and Standard Deviation across folds | |
| metrics_to_agg = ["accuracy", "balanced_accuracy", "auroc", "auprc", "sensitivity", "specificity", "f1", "mcc"] | |
| aggregated = {} | |
| for m in metrics_to_agg: | |
| vals = [f[m] for f in fold_summaries] | |
| aggregated[m] = { | |
| "mean": float(np.mean(vals)), | |
| "std": float(np.std(vals)), | |
| "formatted": f"{np.mean(vals):.4f} ± {np.std(vals):.4f}", | |
| } | |
| # Compute Pooled Out-Of-Fold (OOF) metrics | |
| oof_y = np.array(all_oof_labels) | |
| oof_p = np.array(all_oof_preds) | |
| oof_metrics = compute_metrics(oof_y, oof_p) | |
| cv_results = { | |
| "n_splits": config.n_splits, | |
| "fold_results": fold_summaries, | |
| "mean_std_summary": aggregated, | |
| "pooled_oof_metrics": oof_metrics, | |
| } | |
| # Log and Save CV Summary Report | |
| logging.info("\n" + "=" * 75) | |
| logging.info(f"=== {config.n_splits}-FOLD CROSS-VALIDATION SUMMARY RESULTS ===") | |
| logging.info("-" * 75) | |
| logging.info(f"{'Metric':<25} {'Mean ± Std Across Folds':<30} {'Pooled OOF':<15}") | |
| logging.info("-" * 75) | |
| for m in metrics_to_agg: | |
| logging.info(f"{m:<25} {aggregated[m]['formatted']:<30} {oof_metrics[m]:.4f}") | |
| logging.info("=" * 75 + "\n") | |
| results_dir = os.path.join(config.drive_dir, "results") | |
| os.makedirs(results_dir, exist_ok=True) | |
| with open(os.path.join(results_dir, "5fold_cv_summary.json"), "w") as f: | |
| json.dump(cv_results, f, indent=2) | |
| # Export formatted Markdown summary table | |
| md_lines = [ | |
| f"# {config.n_splits}-Fold Stratified Cross-Validation Results", | |
| "", | |
| f"**Model Backbone:** `{config.backbone}` | **Dataset:** MRNet ACL | **Folds:** {config.n_splits}", | |
| "", | |
| "| Metric | Mean ± Std Across Folds | Pooled Out-of-Fold | Fold 1 | Fold 2 | Fold 3 | Fold 4 | Fold 5 |", | |
| "| :--- | :--- | :--- | :--- | :--- | :--- | :--- | :--- |", | |
| ] | |
| for m in metrics_to_agg: | |
| f_vals = " | ".join([f"{f[m]:.4f}" for f in fold_summaries]) | |
| md_lines.append(f"| **{m.replace('_', ' ').title()}** | `{aggregated[m]['formatted']}` | `{oof_metrics[m]:.4f}` | {f_vals} |") | |
| with open(os.path.join(results_dir, "5fold_cv_summary.md"), "w") as f: | |
| f.write("\n".join(md_lines) + "\n") | |
| return cv_results | |