97 lines
3.1 KiB
Python
97 lines
3.1 KiB
Python
"""Result dataclasses and serialisation helpers for V2 fold outputs."""
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
from typing import Optional
|
|
|
|
import numpy as np
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Primitive helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def _nan() -> float:
|
|
return float("nan")
|
|
|
|
|
|
def _f(v) -> Optional[float]:
|
|
"""Round a scalar to 6 dp, return None for nan/None."""
|
|
if v is None or (isinstance(v, float) and np.isnan(v)):
|
|
return None
|
|
return round(float(v), 6)
|
|
|
|
|
|
def _sv(vec) -> Optional[str]:
|
|
"""Serialise a vector to a pipe-separated string, or None."""
|
|
if vec is None:
|
|
return None
|
|
return "|".join(f"{float(v):.4f}" for v in vec)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Dataclasses
|
|
# ---------------------------------------------------------------------------
|
|
|
|
@dataclass
|
|
class FoldResult:
|
|
mode: str
|
|
fold: int
|
|
# Epoch where each model hit its peak val AUC
|
|
best_epoch_single: int # SingleEyeHT — selected by ensemble val AUC
|
|
best_epoch_bilat: int # BilateralHT — selected by bilateral val AUC
|
|
# Classic (eye-level eval of SingleEyeHT; n = 2 * ensemble_val_n)
|
|
classic_val_auc: float
|
|
classic_val_acc: float
|
|
classic_val_kappa: float
|
|
classic_val_mcc: float
|
|
classic_val_f1: float
|
|
classic_val_recall: Optional[str]
|
|
classic_val_ece: float
|
|
classic_val_threshold: float
|
|
classic_val_bias: Optional[str]
|
|
classic_val_n: int
|
|
# Ensemble (patient-level eval of same SingleEyeHT)
|
|
ensemble_val_auc: float
|
|
ensemble_val_acc: float
|
|
ensemble_val_kappa: float
|
|
ensemble_val_mcc: float
|
|
ensemble_val_f1: float
|
|
ensemble_val_recall: Optional[str]
|
|
ensemble_val_ece: float
|
|
ensemble_val_threshold: float
|
|
ensemble_val_bias: Optional[str]
|
|
ensemble_val_n: int
|
|
# Bilateral (BilateralHT patient-level)
|
|
bilat_val_auc: float
|
|
bilat_val_acc: float
|
|
bilat_val_kappa: float
|
|
bilat_val_mcc: float
|
|
bilat_val_f1: float
|
|
bilat_val_recall: Optional[str]
|
|
bilat_val_ece: float
|
|
bilat_val_threshold: float
|
|
bilat_val_bias: Optional[str]
|
|
bilat_val_n: int
|
|
# Holdout metrics (evaluated at best val epoch; nan if no holdout)
|
|
classic_holdout_auc: float
|
|
classic_holdout_acc: float
|
|
ensemble_holdout_auc: float
|
|
ensemble_holdout_acc: float
|
|
bilat_holdout_auc: float
|
|
bilat_holdout_acc: float
|
|
holdout_n: int # number of holdout bilateral samples
|
|
# Training sample counts
|
|
single_train_n: int
|
|
bilat_train_n: int
|
|
|
|
|
|
@dataclass
|
|
class FoldArtifacts:
|
|
y_true_classic: Optional[np.ndarray]
|
|
probs_classic: Optional[np.ndarray]
|
|
y_true_ensemble: Optional[np.ndarray]
|
|
probs_ensemble: Optional[np.ndarray]
|
|
y_true_bilat: Optional[np.ndarray]
|
|
probs_bilat: Optional[np.ndarray]
|