diff --git a/docs/report/progression.png b/docs/report/progression.png deleted file mode 100644 index 5cd3908..0000000 Binary files a/docs/report/progression.png and /dev/null differ diff --git a/docs/report/test_predictions.png b/docs/report/test_predictions.png deleted file mode 100644 index 5ae35b4..0000000 Binary files a/docs/report/test_predictions.png and /dev/null differ diff --git a/kaggle/run-evaluation.ipynb b/kaggle/run-evaluation.ipynb new file mode 100644 index 0000000..bc18498 --- /dev/null +++ b/kaggle/run-evaluation.ipynb @@ -0,0 +1,151 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": null, + "id": "99852947", + "metadata": {}, + "outputs": [], + "source": [ + "!git clone https://github.com/CfM47/ML-Project.git" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "89e0509f", + "metadata": {}, + "outputs": [], + "source": [ + "import sys\n", + "from pathlib import Path" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2ba6b9c0", + "metadata": {}, + "outputs": [], + "source": [ + "!ls /kaggle/input/training-notebook-output" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "05fff357", + "metadata": {}, + "outputs": [], + "source": [ + "segmentation_root = Path('/kaggle/input/segmentations-images-automl/pictures')\n", + "\n", + "TRAIN_LABELED = segmentation_root / 'vega_3_tescan_labeled_images'\n", + "TRAIN_UNLABELED = segmentation_root / 'vega_3_tescan_unlabeled_images'\n", + "TEST_UNLABELED = segmentation_root / 'sampled_unlabeled'\n", + "TEST_LABELED = segmentation_root / 'sampled_labeled'\n", + "\n", + "# Pretrained model from previous training run\n", + "PRETRAINED_MODEL = Path('/kaggle/input/training-notebook-output/swin_results/model.pt')\n", + "\n", + "WORKING_DIR = Path('/kaggle/working')\n", + "PROJECT_ROOT = WORKING_DIR / 'ML-Project'\n", + "OUTPUT_DIR = WORKING_DIR / 'swin_results'" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "bbcc6efc", + "metadata": {}, + "outputs": [], + "source": [ + "if str(PROJECT_ROOT) not in sys.path:\n", + " sys.path.insert(0, str(PROJECT_ROOT))" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "bd42e98a", + "metadata": {}, + "outputs": [], + "source": [ + "from model.swin.config import SwinTrainingConfig\n", + "\n", + "config = SwinTrainingConfig(\n", + " output_dir=OUTPUT_DIR,\n", + " device='auto',\n", + ")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "39900be6", + "metadata": {}, + "outputs": [], + "source": [ + "from model.swin.train import run_evaluation_only\n", + "\n", + "result = run_evaluation_only(\n", + " TRAIN_UNLABELED,\n", + " TRAIN_LABELED,\n", + " TEST_UNLABELED,\n", + " TEST_LABELED,\n", + " PRETRAINED_MODEL,\n", + " config=config,\n", + ")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "3d554aad", + "metadata": {}, + "outputs": [], + "source": [ + "result.predictions_figure" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "3a5d834b", + "metadata": {}, + "outputs": [], + "source": [ + "print(\"Full Test Set Metrics:\")\n", + "for metric_name, value in result.test_metrics.full_metrics.items():\n", + " print(f\" {metric_name}: {value:.4f}\")\n", + "\n", + "print(\"\\nLeave-Two-Out Summary:\")\n", + "print(f\" Subsets: {result.test_metrics.num_subsets} of size {result.test_metrics.subset_size}\")\n", + "for metric_name in sorted(result.test_metrics.subset_means.keys()):\n", + " mean_val = result.test_metrics.subset_means[metric_name]\n", + " std_val = result.test_metrics.subset_stds[metric_name]\n", + " print(f\" {metric_name}: {mean_val:.4f} ± {std_val:.4f}\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "e5f6g7h8", + "metadata": {}, + "outputs": [], + "source": [ + "# Display all histograms\n", + "for metric_name, fig in sorted(result.histograms.items()):\n", + " print(f'\\n{metric_name}:')\n", + " display(fig)" + ] + } + ], + "metadata": { + "language_info": { + "name": "python" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/kaggle/run-training.ipynb b/kaggle/run-training.ipynb index e17bca4..81330f2 100644 --- a/kaggle/run-training.ipynb +++ b/kaggle/run-training.ipynb @@ -85,7 +85,7 @@ "source": [ "from model.swin.train import run_final_training\n", "\n", - "model, test_metrics, mask_pairs, predictions_fig, loss_curves_fig = run_final_training(\n", + "result = run_final_training(\n", " TRAIN_UNLABELED,\n", " TRAIN_LABELED,\n", " TEST_UNLABELED,\n", @@ -101,7 +101,7 @@ "metadata": {}, "outputs": [], "source": [ - "predictions_fig" + "result.predictions_figure" ] }, { @@ -111,8 +111,8 @@ "metadata": {}, "outputs": [], "source": [ - "if loss_curves_fig is not None:\n", - " display(loss_curves_fig)\n", + "if result.loss_curves_figure is not None:\n", + " display(result.loss_curves_figure)\n", "else:\n", " print('Loss curves not available (training history validation failed)')" ] @@ -124,9 +124,29 @@ "metadata": {}, "outputs": [], "source": [ - "print(\"Test Metrics:\")\n", - "for metric_name, value in test_metrics.items():\n", - " print(f\" {metric_name}: {value:.4f}\")" + "print(\"Full Test Set Metrics:\")\n", + "for metric_name, value in result.test_metrics.full_metrics.items():\n", + " print(f\" {metric_name}: {value:.4f}\")\n", + "\n", + "print(\"\\nLeave-Two-Out Summary:\")\n", + "print(f\" Subsets: {result.test_metrics.num_subsets} of size {result.test_metrics.subset_size}\")\n", + "for metric_name in sorted(result.test_metrics.subset_means.keys()):\n", + " mean_val = result.test_metrics.subset_means[metric_name]\n", + " std_val = result.test_metrics.subset_stds[metric_name]\n", + " print(f\" {metric_name}: {mean_val:.4f} ± {std_val:.4f}\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "e5f6g7h8", + "metadata": {}, + "outputs": [], + "source": [ + "# Display all histograms\n", + "for metric_name, fig in sorted(result.histograms.items()):\n", + " print(f'\\n{metric_name}:')\n", + " display(fig)" ] } ], diff --git a/model/swin/__init__.py b/model/swin/__init__.py index 904030e..9651cf8 100644 --- a/model/swin/__init__.py +++ b/model/swin/__init__.py @@ -1,10 +1,15 @@ """Swin model training and validation submodule.""" from model.swin.config import SwinTrainingConfig -from model.swin.train import run_final_training, run_percentage_validation +from model.swin.train import ( + run_evaluation_only, + run_final_training, + run_percentage_validation, +) __all__ = [ "SwinTrainingConfig", "run_percentage_validation", "run_final_training", + "run_evaluation_only", ] diff --git a/model/swin/evaluation.py b/model/swin/evaluation.py index 6e2ecc4..0d71c2a 100644 --- a/model/swin/evaluation.py +++ b/model/swin/evaluation.py @@ -4,28 +4,57 @@ from auto_ml.implementations.evaluators import ( AccuracyEvaluator, - DiceMacroAverageEvaluator, + AutoencoderMaskEvaluator, + DiceClass0Evaluator, + DiceClass1Evaluator, + IoUClass0Evaluator, + IoUClass1Evaluator, + PrecisionClass0Evaluator, + PrecisionClass1Evaluator, + RecallClass0Evaluator, + RecallClass1Evaluator, ) from auto_ml.implementations.nodes import EvaluatorNode from auto_ml.implementations.segmentators.swin import SwinModel from auto_ml.interfaces import MaskPair, SegmentationDatasetInterface -def create_evaluator() -> EvaluatorNode: +def create_evaluator(dataset: SegmentationDatasetInterface) -> EvaluatorNode: """ - Create evaluator with Dice (F1) and Accuracy metrics only. + Create evaluator with all available metrics for binary segmentation. + + Args: + dataset: Dataset used for training the Mask_Cohesion autoencoder evaluator. Returns: - Configured EvaluatorNode. + Configured EvaluatorNode with all metrics. """ - return EvaluatorNode( - evaluators={ - "Dice_Macro": DiceMacroAverageEvaluator(), - "Accuracy": AccuracyEvaluator(), - }, - name="SwinValidationEvaluator", - ) + evaluators = { + # General + "Accuracy": AccuracyEvaluator(), + # Autoencoder (requires training on reference masks) + "Mask_Cohesion": AutoencoderMaskEvaluator( + reference_masks=dataset.masks, + latent_dim=8, + epochs=40, + nu=0.5, + device="auto", + ), + # IoU Metrics (binary: Class0 and Class1 only) + "IoU_Class0": IoUClass0Evaluator(), + "IoU_Class1": IoUClass1Evaluator(), + # Dice Metrics (binary: Class0 and Class1 only) + "Dice_Class0": DiceClass0Evaluator(), + "Dice_Class1": DiceClass1Evaluator(), + # Precision Metrics (binary: Class0 and Class1 only) + "Precision_Class0": PrecisionClass0Evaluator(), + "Precision_Class1": PrecisionClass1Evaluator(), + # Recall Metrics (binary: Class0 and Class1 only) + "Recall_Class0": RecallClass0Evaluator(), + "Recall_Class1": RecallClass1Evaluator(), + } + return EvaluatorNode(evaluators=evaluators, name="SwinValidationEvaluator") def evaluate_model( @@ -56,13 +85,41 @@ def evaluate_model( metrics: Dict[str, float] = {} for metric_name, values in evaluation_results.items(): if isinstance(values, list) and len(values) > 0: - metrics[metric_name] = values[0] + metrics[metric_name] = float(values[0]) else: metrics[metric_name] = float(values) + # Compute F1 scores from Precision and Recall + _add_f1_metrics(metrics) + return metrics, mask_pairs +def _add_f1_metrics(metrics: Dict[str, float]) -> None: + """ + Compute F1 scores from Precision and Recall and add them to metrics dict. + + Formula: F1 = 2 * (Precision * Recall) / (Precision + Recall) + + Args: + metrics: Metrics dictionary to update in-place. + + """ + for class_id in [0, 1]: + precision_key = f"Precision_Class{class_id}" + recall_key = f"Recall_Class{class_id}" + f1_key = f"F1_Class{class_id}" + + if precision_key in metrics and recall_key in metrics: + precision = metrics[precision_key] + recall = metrics[recall_key] + denominator = precision + recall + if denominator > 0: + metrics[f1_key] = 2 * (precision * recall) / denominator + else: + metrics[f1_key] = 0.0 + + def extract_metrics_from_evaluation( evaluation_results: Dict[str, Any], ) -> Dict[str, float]: @@ -79,7 +136,71 @@ def extract_metrics_from_evaluation( metrics: Dict[str, float] = {} for metric_name, values in evaluation_results.items(): if isinstance(values, list) and len(values) > 0: - metrics[metric_name] = values[0] + metrics[metric_name] = float(values[0]) else: metrics[metric_name] = float(values) return metrics + + +def evaluate_leave_two_out( + mask_pairs: List[MaskPair], + evaluator: EvaluatorNode, +) -> Tuple[Dict[str, float], Dict[str, float], Dict[str, List[float]]]: + """ + Evaluate all leave-two-out subsets for robustness analysis. + + For n samples, evaluate all C(n,2) subsets of size n-2. + + Args: + mask_pairs: List of (predicted_mask, real_mask) tuples from full test set. + evaluator: Configured EvaluatorNode for computing metrics. + + Returns: + Tuple of (subset_means, subset_stds, subset_distributions). + + """ + import itertools + + import numpy as np + + n = len(mask_pairs) + subset_size = n - 2 + num_subsets = n * (n - 1) // 2 # C(n, 2) + + print( + f"Leave-two-out evaluation: {num_subsets} subsets of size " + f"{subset_size} from {n} samples", + ) + + # Generate all combinations of 2 indices to exclude + exclude_pairs = list(itertools.combinations(range(n), 2)) + + # Initialize distributions dict + distributions: Dict[str, List[float]] = {} + + for exclude_idx_pair in exclude_pairs: + # Create subset by excluding two samples + subset = [mask_pairs[i] for i in range(n) if i not in exclude_idx_pair] + + # Evaluate subset + evaluation_results = evaluator.evaluate([subset]) + + # Extract metrics + subset_metrics = extract_metrics_from_evaluation(evaluation_results) + _add_f1_metrics(subset_metrics) + + # Add to distributions + for metric_name, value in subset_metrics.items(): + if metric_name not in distributions: + distributions[metric_name] = [] + distributions[metric_name].append(value) + + # Compute means and stds + subset_means: Dict[str, float] = {} + subset_stds: Dict[str, float] = {} + + for metric_name, values in distributions.items(): + subset_means[metric_name] = float(np.mean(values)) + subset_stds[metric_name] = float(np.std(values)) + + return subset_means, subset_stds, distributions diff --git a/model/swin/metrics.py b/model/swin/metrics.py index 5b893fe..8bd58c7 100644 --- a/model/swin/metrics.py +++ b/model/swin/metrics.py @@ -1,10 +1,16 @@ """Metrics dataclasses for tracking training and validation results.""" from dataclasses import dataclass, field -from typing import Dict, List +from typing import TYPE_CHECKING, Any, Dict, List import numpy as np +if TYPE_CHECKING: + from matplotlib.figure import Figure + + from auto_ml.implementations.segmentators.swin import SwinModel + from auto_ml.interfaces import MaskPair + @dataclass class TrainingHistory: @@ -68,9 +74,34 @@ class FoldMetrics: fold: int train_history: List[Dict[str, float]] = field(default_factory=list) - # Final metrics from evaluator - dice_macro: float = 0.0 - accuracy: float = 0.0 + # All metrics from evaluator + metrics: Dict[str, float] = field(default_factory=dict) + + @property + def dice_macro(self) -> float: + """ + Return Dice macro metric for backward compatibility. + + Raises: + KeyError: If Dice_Macro metric not found. + + """ + if "Dice_Macro" not in self.metrics: + raise KeyError("Metric 'Dice_Macro' not found in fold metrics") + return self.metrics["Dice_Macro"] + + @property + def accuracy(self) -> float: + """ + Return Accuracy metric for backward compatibility. + + Raises: + KeyError: If Accuracy metric not found. + + """ + if "Accuracy" not in self.metrics: + raise KeyError("Metric 'Accuracy' not found in fold metrics") + return self.metrics["Accuracy"] @property def final_train_loss(self) -> float: @@ -104,37 +135,96 @@ class PercentageMetrics: percentage: int fold_metrics: List[FoldMetrics] = field(default_factory=list) - # --- Dice Macro (F1) --- + # --- Dynamic Metric Access --- + + def get_metric_mean(self, metric_name: str) -> float: + """ + Calculate mean of a specific metric across folds. + + Args: + metric_name: Name of the metric (e.g., "Dice_Macro", "IoU_Class1"). + + Returns: + Mean value across folds. + + Raises: + ValueError: If no fold metrics available. + KeyError: If metric not found in any fold. + + """ + if not self.fold_metrics: + raise ValueError("No fold metrics available") + for fm in self.fold_metrics: + if metric_name not in fm.metrics: + raise KeyError( + f"Metric '{metric_name}' not found in fold {fm.fold} metrics", + ) + values = [fm.metrics[metric_name] for fm in self.fold_metrics] + return float(np.mean(values)) + + def get_metric_std(self, metric_name: str) -> float: + """ + Calculate std of a specific metric across folds. + + Args: + metric_name: Name of the metric (e.g., "Dice_Macro", "IoU_Class1"). + + Returns: + Standard deviation across folds. + + Raises: + ValueError: If no fold metrics available. + KeyError: If metric not found in any fold. + + """ + if not self.fold_metrics: + raise ValueError("No fold metrics available") + for fm in self.fold_metrics: + if metric_name not in fm.metrics: + raise KeyError( + f"Metric '{metric_name}' not found in fold {fm.fold} metrics", + ) + values = [fm.metrics[metric_name] for fm in self.fold_metrics] + return float(np.std(values)) + + def get_all_metric_names(self) -> List[str]: + """ + Return all metric names available in the fold metrics. + + Returns: + List of metric names. + + Raises: + ValueError: If no fold metrics available. + + """ + if not self.fold_metrics: + raise ValueError("No fold metrics available") + return list(self.fold_metrics[0].metrics.keys()) + + # --- Dice Macro (F1) - Backward Compatibility --- @property def mean_dice_macro(self) -> float: """Calculate mean Dice macro across folds.""" - if not self.fold_metrics: - return 0.0 - return float(np.mean([fm.dice_macro for fm in self.fold_metrics])) + return self.get_metric_mean("Dice_Macro") @property def std_dice_macro(self) -> float: """Calculate std of Dice macro across folds.""" - if not self.fold_metrics: - return 0.0 - return float(np.std([fm.dice_macro for fm in self.fold_metrics])) + return self.get_metric_std("Dice_Macro") - # --- Accuracy --- + # --- Accuracy - Backward Compatibility --- @property def mean_accuracy(self) -> float: """Calculate mean accuracy across folds.""" - if not self.fold_metrics: - return 0.0 - return float(np.mean([fm.accuracy for fm in self.fold_metrics])) + return self.get_metric_mean("Accuracy") @property def std_accuracy(self) -> float: """Calculate std of accuracy across folds.""" - if not self.fold_metrics: - return 0.0 - return float(np.std([fm.accuracy for fm in self.fold_metrics])) + return self.get_metric_std("Accuracy") # --- Training Loss --- @@ -167,3 +257,52 @@ def std_final_val_loss(self) -> float: if not self.fold_metrics: return 0.0 return float(np.std([fm.final_val_loss for fm in self.fold_metrics])) + + +@dataclass +class SubsetMetrics: + """Store leave-two-out evaluation metrics.""" + + subset_size: int + num_subsets: int + + # Full test set metrics: {"Accuracy": 0.85, ...} + full_metrics: Dict[str, float] + + # Mean across subsets: {"Accuracy": 0.84, ...} + subset_means: Dict[str, float] + + # Std across subsets: {"Accuracy": 0.02, ...} + subset_stds: Dict[str, float] + + # Full distributions: {"Accuracy": [0.84, 0.86, ...], ...} + subset_distributions: Dict[str, List[float]] + + def to_dict(self) -> Dict[str, Any]: + """Convert to JSON-serializable dictionary.""" + return { + "subset_size": self.subset_size, + "num_subsets": self.num_subsets, + "full_metrics": self.full_metrics, + "subset_means": self.subset_means, + "subset_stds": self.subset_stds, + "subset_distributions": self.subset_distributions, + } + + +@dataclass +class TrainingResult: + """Store all outputs from final training.""" + + model: "SwinModel" + test_metrics: SubsetMetrics + mask_pairs: List["MaskPair"] + predictions_figure: "Figure" + loss_curves_figure: "Figure | None" + histograms: Dict[str, "Figure"] + + def to_dict(self) -> Dict[str, Any]: + """Convert to JSON-serializable dictionary (excludes non-serializable objects).""" # noqa: E501 + return { + "test_metrics": self.test_metrics.to_dict(), + } diff --git a/model/swin/train.py b/model/swin/train.py index 545eee3..60fd7c5 100644 --- a/model/swin/train.py +++ b/model/swin/train.py @@ -1,11 +1,13 @@ """ Swin Segmentation Training and Validation Module. -Provide two entry points: +Provide three entry points: 1. run_percentage_validation: K-fold cross-validation with varying training percentages 2. run_final_training: Train on 100% data and evaluate on test set +3. run_evaluation_only: Load pretrained model and evaluate on test set """ +import json from pathlib import Path from typing import Dict, List, Tuple @@ -17,9 +19,20 @@ from auto_ml.interfaces import MaskPair, SegmentationDatasetInterface from model.swin.config import SwinTrainingConfig from model.swin.data import create_augmentator, create_kfold_splits, subsample_dataset -from model.swin.evaluation import create_evaluator, evaluate_model -from model.swin.metrics import FoldMetrics, PercentageMetrics, TrainingHistory +from model.swin.evaluation import ( + create_evaluator, + evaluate_leave_two_out, + evaluate_model, +) +from model.swin.metrics import ( + FoldMetrics, + PercentageMetrics, + SubsetMetrics, + TrainingHistory, + TrainingResult, +) from model.swin.visualization import ( + plot_metric_histograms, plot_progression_grid, plot_results, plot_training_loss_curves, @@ -124,11 +137,12 @@ def run_final_training( test_unlabeled_dir: str | Path, test_labeled_dir: str | Path, config: SwinTrainingConfig | None = None, -) -> Tuple[SwinModel, Dict[str, float], List[MaskPair], Figure, Figure | None]: +) -> TrainingResult: """ Train final model on 100% data and evaluate on test set. - Save model weights and prediction visualizations to output directory. + Save model weights, prediction visualizations, metric histograms, and + results JSON to output directory. Args: train_unlabeled_dir: Directory with unlabeled training images. @@ -138,8 +152,7 @@ def run_final_training( config: Training configuration. Uses defaults if None. Returns: - Tuple of (trained_model, test_metrics, test_mask_pairs, predictions Figure, - loss_curves Figure or None if history validation failed). + TrainingResult containing model, metrics, mask_pairs, and figures. """ if config is None: @@ -166,32 +179,78 @@ def run_final_training( # Train final model model, training_history = _train_final_model(train_dataset, config) - # Evaluate on test set - test_metrics, test_mask_pairs = _evaluate_on_test(model, test_dataset) - # Save model model_path = config.output_dir / "model.pt" _save_model(model, model_path) - # Visualize predictions - viz_path = config.output_dir / "test_predictions.png" - predictions_fig = visualize_predictions( + # Evaluate model and save results + return _evaluate_model( + model, + train_dataset, test_dataset, - test_mask_pairs, - num_samples=config.num_test_visualizations, - output_path=viz_path, + config, + training_history, ) - # Plot training loss curves if history is valid - loss_curves_fig: Figure | None = None - if training_history is not None: - loss_curves_path = config.output_dir / "training_loss_curves.png" - loss_curves_fig = plot_training_loss_curves( - training_history, - output_path=loss_curves_path, - ) - return model, test_metrics, test_mask_pairs, predictions_fig, loss_curves_fig +def run_evaluation_only( + train_unlabeled_dir: str | Path, + train_labeled_dir: str | Path, + test_unlabeled_dir: str | Path, + test_labeled_dir: str | Path, + pretrained_model_path: str | Path, + config: SwinTrainingConfig | None = None, +) -> TrainingResult: + """ + Load a pretrained model and evaluate on test set. + + Entry point for evaluating without training. Perform leave-two-out analysis + and save prediction visualizations, metric histograms, and results JSON. + + Args: + train_unlabeled_dir: Directory with unlabeled training images. + train_labeled_dir: Directory with labeled training masks. + test_unlabeled_dir: Directory with unlabeled test images. + test_labeled_dir: Directory with labeled test masks. + pretrained_model_path: Path to pretrained model weights (.pt file). + config: Training configuration. Uses defaults if None. + + Returns: + TrainingResult containing model, metrics, mask_pairs, and figures. + + """ + if config is None: + config = SwinTrainingConfig() + + # Setup output directory + config.output_dir.mkdir(parents=True, exist_ok=True) + + # Load datasets + print("Loading training dataset...") + train_dataset = load_dataset_from_directories( + Path(train_unlabeled_dir), + Path(train_labeled_dir), + ) + print(f"Loaded {len(train_dataset)} training samples") + + print("Loading test dataset...") + test_dataset = load_dataset_from_directories( + Path(test_unlabeled_dir), + Path(test_labeled_dir), + ) + print(f"Loaded {len(test_dataset)} test samples") + + # Load pretrained model + model = _load_model(Path(pretrained_model_path), config) + + # Evaluate model (no training history since we loaded a pretrained model) + return _evaluate_model( + model, + train_dataset, + test_dataset, + config, + training_history=None, + ) # ============================================================================== @@ -221,14 +280,10 @@ def _run_validation_loop( best_models_by_percentage[percentage] = best_model print(f"\n {percentage}% Summary:") - print( - f" Mean F1 (Dice): {pct_metrics.mean_dice_macro:.4f} " - f"± {pct_metrics.std_dice_macro:.4f}", - ) - print( - f" Mean Accuracy: {pct_metrics.mean_accuracy:.4f} " - f"± {pct_metrics.std_accuracy:.4f}", - ) + for metric_name in pct_metrics.get_all_metric_names(): + mean_val = pct_metrics.get_metric_mean(metric_name) + std_val = pct_metrics.get_metric_std(metric_name) + print(f" {metric_name}: {mean_val:.4f} ± {std_val:.4f}") return all_metrics, best_models_by_percentage @@ -262,10 +317,9 @@ def _run_percentage_experiment( best_dice = fold_metrics.dice_macro best_model = model - print( - f" F1 (Dice): {fold_metrics.dice_macro:.4f}, " - f"Accuracy: {fold_metrics.accuracy:.4f}", - ) + print(f" Fold {fold + 1} Metrics:") + for metric_name, value in sorted(fold_metrics.metrics.items()): + print(f" {metric_name}: {value:.4f}") # best_model is guaranteed to be set since we have at least one fold assert best_model is not None @@ -299,15 +353,14 @@ def _train_fold( # Train model train_result = model.train(aug_train, validation_dataset=val_dataset) - # Evaluate on validation set - evaluator = create_evaluator() + # Evaluate on validation set (autoencoder trained on training set) + evaluator = create_evaluator(train_dataset) metrics, _ = evaluate_model(model, val_dataset, evaluator) fold_metrics = FoldMetrics( fold=fold, train_history=train_result.history, - dice_macro=metrics.get("Dice_Macro", 0.0), - accuracy=metrics.get("Accuracy", 0.0), + metrics=metrics, ) return fold_metrics, model @@ -353,19 +406,129 @@ def _train_final_model( def _evaluate_on_test( model: SwinModel, test_dataset: SegmentationDatasetInterface, -) -> Tuple[Dict[str, float], List[MaskPair]]: - """Evaluate trained model on test dataset (no augmentation).""" + train_dataset: SegmentationDatasetInterface, +) -> Tuple[SubsetMetrics, List[MaskPair]]: + """ + Evaluate trained model on test dataset with leave-two-out analysis. + + Args: + model: Trained SwinModel. + test_dataset: Test dataset to evaluate on. + train_dataset: Training dataset for Mask_Cohesion autoencoder reference. + + Returns: + Tuple of (SubsetMetrics, mask_pairs). + + """ print("\n" + "=" * 60) print("Evaluating on test set") print("=" * 60) - evaluator = create_evaluator() - metrics, mask_pairs = evaluate_model(model, test_dataset, evaluator) + # Create evaluator (trained on training set for Mask_Cohesion) + evaluator = create_evaluator(train_dataset) + + # Evaluate full test set + full_metrics, mask_pairs = evaluate_model(model, test_dataset, evaluator) + + print("\nFull Test Set Metrics:") + for metric_name, value in sorted(full_metrics.items()): + print(f" {metric_name}: {value:.4f}") + + # Perform leave-two-out evaluation + subset_means, subset_stds, subset_distributions = evaluate_leave_two_out( + mask_pairs, + evaluator, + ) + + # Build SubsetMetrics + n = len(mask_pairs) + subset_metrics = SubsetMetrics( + subset_size=n - 2, + num_subsets=n * (n - 1) // 2, + full_metrics=full_metrics, + subset_means=subset_means, + subset_stds=subset_stds, + subset_distributions=subset_distributions, + ) - print(f"Test F1 (Dice): {metrics.get('Dice_Macro', 0.0):.4f}") - print(f"Test Accuracy: {metrics.get('Accuracy', 0.0):.4f}") + print("\nLeave-Two-Out Summary:") + for metric_name in sorted(subset_means.keys()): + mean_val = subset_means[metric_name] + std_val = subset_stds[metric_name] + print(f" {metric_name}: {mean_val:.4f} ± {std_val:.4f}") - return metrics, mask_pairs + return subset_metrics, mask_pairs + + +def _evaluate_model( + model: SwinModel, + train_dataset: SegmentationDatasetInterface, + test_dataset: SegmentationDatasetInterface, + config: SwinTrainingConfig, + training_history: TrainingHistory | None = None, +) -> TrainingResult: + """ + Evaluate a model on test set with leave-two-out analysis. + + Save predictions, histograms, and results JSON to output directory. + + Args: + model: Trained or loaded SwinModel. + train_dataset: Training dataset for Mask_Cohesion autoencoder reference. + test_dataset: Test dataset to evaluate on. + config: Training configuration with output directory. + training_history: Optional training history for loss curves plot. + + Returns: + TrainingResult containing model, metrics, mask_pairs, and figures. + + """ + # Evaluate on test set with leave-two-out analysis + test_metrics, test_mask_pairs = _evaluate_on_test( + model, + test_dataset, + train_dataset, + ) + + # Visualize predictions + viz_path = config.output_dir / "test_predictions.png" + predictions_fig = visualize_predictions( + test_dataset, + test_mask_pairs, + num_samples=config.num_test_visualizations, + output_path=viz_path, + ) + + # Plot training loss curves if history is valid + loss_curves_fig: Figure | None = None + if training_history is not None: + loss_curves_path = config.output_dir / "training_loss_curves.png" + loss_curves_fig = plot_training_loss_curves( + training_history, + output_path=loss_curves_path, + ) + + # Plot metric histograms + histograms = plot_metric_histograms( + test_metrics.subset_distributions, + test_metrics.full_metrics, + output_dir=config.output_dir, + ) + + # Build result object + result = TrainingResult( + model=model, + test_metrics=test_metrics, + mask_pairs=test_mask_pairs, + predictions_figure=predictions_fig, + loss_curves_figure=loss_curves_fig, + histograms=histograms, + ) + + # Save results JSON + _save_results_json(result, config.output_dir / "results.json") + + return result # ============================================================================== @@ -393,6 +556,22 @@ def _save_model(model: SwinModel, path: Path) -> None: print(f"Saved model to {path}") +def _save_results_json(result: TrainingResult, path: Path) -> None: + """Save training results to JSON file.""" + results_dict = result.to_dict() + with open(path, "w") as f: + json.dump(results_dict, f, indent=2) + print(f"Saved results to {path}") + + +def _load_model(path: Path, config: SwinTrainingConfig) -> SwinModel: + """Load model weights from disk.""" + model = _create_swin_model(config) + model.model.load_state_dict(torch.load(path, weights_only=True)) + print(f"Loaded model from {path}") + return model + + def _collect_progression_predictions( best_models: Dict[int, SwinModel], test_dataset: SegmentationDatasetInterface, diff --git a/model/swin/visualization.py b/model/swin/visualization.py index 5057b6e..9bc6704 100644 --- a/model/swin/visualization.py +++ b/model/swin/visualization.py @@ -416,3 +416,78 @@ def _plot_loss_curves( ax.set_title(title) ax.legend() ax.grid(True, alpha=0.3) + + +def plot_metric_histograms( + subset_distributions: Dict[str, List[float]], + full_metrics: Dict[str, float], + output_dir: Path | None = None, +) -> Dict[str, Figure]: + """ + Create individual histogram for each metric showing leave-two-out distribution. + + Each histogram shows the distribution of metric values across all subsets, + with a vertical line indicating the full test set metric value. + + Args: + subset_distributions: Dict mapping metric name to list of subset values. + full_metrics: Dict mapping metric name to full test set value. + output_dir: Optional directory to save histogram files. + + Returns: + Dict mapping metric name to Figure. + + """ + histograms: Dict[str, Figure] = {} + + for metric_name, values in subset_distributions.items(): + fig, ax = plt.subplots(figsize=(8, 5)) + + # Plot histogram + ax.hist( + values, + bins=20, + color="steelblue", + edgecolor="white", + alpha=0.7, + ) + + # Add vertical line for full test set value + if metric_name in full_metrics: + full_value = full_metrics[metric_name] + ax.axvline( + full_value, + color="red", + linestyle="--", + linewidth=2, + label=f"Full test set: {full_value:.4f}", + ) + + # Calculate and display statistics + mean_val = float(np.mean(values)) + std_val = float(np.std(values)) + ax.axvline( + mean_val, + color="green", + linestyle="-", + linewidth=2, + label=f"Mean: {mean_val:.4f} (std: {std_val:.4f})", + ) + + ax.set_xlabel(metric_name) + ax.set_ylabel("Frequency") + ax.set_title(f"{metric_name} Distribution (Leave-Two-Out)") + ax.legend(loc="upper right") + ax.grid(True, alpha=0.3) + + plt.tight_layout() + + # Save if output directory provided + if output_dir: + output_path = output_dir / f"histogram_{metric_name}.png" + fig.savefig(output_path, dpi=150, bbox_inches="tight") + print(f"Saved histogram to {output_path}") + + histograms[metric_name] = fig + + return histograms diff --git a/pyproject.toml b/pyproject.toml index 33f31fa..cf97440 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -47,7 +47,7 @@ lint.select = [ ] lint.ignore = ["D100", "D104", "D203", "D212"] -exclude = [".venv", "build", "dist", "__pycache__", ".git", "plots/", "results/"] +exclude = [".venv", "build", "dist", "__pycache__", ".git", "plots/", "results/", '*.ipynb'] format.quote-style = "double" format.indent-style = "space" line-length = 88