Spaces:
Runtime error
Runtime error
| """segroc Streamlit dashboard — v2. | |
| Three evaluation modes | |
| --------------------- | |
| 1. MedSegBench Robustness Cards — browse 5 MedSegBench datasets × 3 | |
| architectures × 3 training seeds from pre-built | |
| artifact bundles. Robustness card is downloadable | |
| as JSON; the page can be printed / saved as PDF | |
| via the browser. | |
| 2. Compare Models — pick two precomputed models on the same dataset | |
| and view their robustness cards side-by-side with | |
| per-metric win/loss indicators. | |
| 3. Upload Your Model — supply a MONAI-compatible .pt checkpoint and a | |
| single input image for a live single-sample | |
| robustness sweep. | |
| Local dev | |
| --------- | |
| uv run streamlit run app.py | |
| """ | |
| from __future__ import annotations | |
| import io | |
| import tempfile | |
| from pathlib import Path | |
| from segroc.data.datasets import ( | |
| list_medsegbench_datasets, | |
| load_medsegbench_split, | |
| ) | |
| import numpy as np | |
| from segroc.model.reporting import ( | |
| MetricSnapshot, | |
| SafetyThresholds, | |
| SafetyAssessment, | |
| assess_safety_gates, | |
| build_model_robustness_card, | |
| ) | |
| from segroc.utils.model_card_bundle import ( | |
| RobustnessBundle, | |
| bundle_from_zip_bytes, | |
| bundle_to_streamlit_eval_results, | |
| bundle_to_zip_bytes, | |
| ) | |
| from segroc.utils.robustness_compare import ( | |
| compare_robustness_bundles, | |
| ) | |
| import pandas as pd | |
| from segroc.data.protocols import list_medsegbench_quickrun_tasks | |
| import streamlit as st | |
| import streamlit.components.v1 as _stc | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| # --------------------------------------------------------------------------- | |
| # Page config (must be first Streamlit call) | |
| # --------------------------------------------------------------------------- | |
| st.set_page_config( | |
| page_title="segroc — Segmentation Robustness Explorer", | |
| layout="wide", | |
| initial_sidebar_state="expanded", | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Global CSS | |
| # --------------------------------------------------------------------------- | |
| st.markdown( | |
| """ | |
| <style> | |
| /* ── Sidebar ──────────────────────────────────────────────── */ | |
| section[data-testid="stSidebar"] { | |
| background: linear-gradient(180deg, #0f172a 0%, #1e293b 100%); | |
| } | |
| section[data-testid="stSidebar"] * { color: #e2e8f0 !important; } | |
| section[data-testid="stSidebar"] .stButton > button { | |
| background: #0ea5e9; | |
| color: #fff !important; | |
| border: none; | |
| border-radius: 8px; | |
| font-weight: 600; | |
| letter-spacing: 0.02em; | |
| } | |
| section[data-testid="stSidebar"] .stButton > button:hover { | |
| background: #0284c7; | |
| } | |
| section[data-testid="stSidebar"] .stSelectbox label, | |
| section[data-testid="stSidebar"] .stRadio label, | |
| section[data-testid="stSidebar"] .stSlider label { | |
| color: #94a3b8 !important; | |
| font-size: 0.8rem; | |
| text-transform: uppercase; | |
| letter-spacing: 0.06em; | |
| } | |
| /* ── Metric cards ─────────────────────────────────────────── */ | |
| div[data-testid="metric-container"] { | |
| background: #f8fafc; | |
| border: 1px solid #e2e8f0; | |
| border-radius: 10px; | |
| padding: 14px 16px; | |
| } | |
| div[data-testid="metric-container"] label { | |
| font-size: 0.78rem; | |
| color: #64748b !important; | |
| text-transform: uppercase; | |
| letter-spacing: 0.05em; | |
| } | |
| div[data-testid="metric-container"] div[data-testid="stMetricValue"] { | |
| font-size: 1.45rem; | |
| font-weight: 700; | |
| color: #0f172a; | |
| } | |
| /* ── Tabs ─────────────────────────────────────────────────── */ | |
| .stTabs [data-baseweb="tab-list"] { | |
| gap: 4px; | |
| border-bottom: 2px solid #e2e8f0; | |
| } | |
| .stTabs [data-baseweb="tab"] { | |
| border-radius: 6px 6px 0 0; | |
| padding: 8px 18px; | |
| font-weight: 500; | |
| color: #64748b; | |
| } | |
| .stTabs [aria-selected="true"] { | |
| background: #f0f9ff; | |
| color: #0ea5e9 !important; | |
| border-bottom: 2px solid #0ea5e9; | |
| } | |
| /* ── Dataframe ────────────────────────────────────────────── */ | |
| .dataframe thead th { | |
| background: #f1f5f9; | |
| font-size: 0.78rem; | |
| text-transform: uppercase; | |
| letter-spacing: 0.04em; | |
| color: #475569; | |
| } | |
| /* ── Callout blocks ───────────────────────────────────────── */ | |
| .segroc-note { | |
| border-left: 4px solid #0ea5a4; | |
| background: #f0fdfa; | |
| color: #0f172a; | |
| border-radius: 8px; | |
| padding: 10px 14px; | |
| margin-bottom: 10px; | |
| font-size: 0.93rem; | |
| } | |
| .segroc-disclaimer { | |
| border-left: 4px solid #f59e0b; | |
| background: #fffbeb; | |
| color: #78350f; | |
| border-radius: 8px; | |
| padding: 10px 14px; | |
| margin-bottom: 12px; | |
| font-size: 0.91rem; | |
| } | |
| .segroc-hero { | |
| background: linear-gradient(135deg, #0f172a 0%, #1e293b 40%, #134e4a 100%); | |
| border-radius: 14px; | |
| padding: 18px 22px; | |
| color: #e2e8f0; | |
| margin-bottom: 14px; | |
| border: 1px solid rgba(148,163,184,0.22); | |
| } | |
| .segroc-hero h2 { margin: 0 0 4px 0; font-weight: 800; letter-spacing: -0.5px; } | |
| .segroc-hero p { margin: 0; color: #94a3b8; font-size: 0.95rem; } | |
| /* ── Comparison indicator badges ──────────────────────────── */ | |
| .badge-better { | |
| display: inline-block; | |
| background: #dcfce7; | |
| color: #166534; | |
| border: 1px solid #86efac; | |
| border-radius: 6px; | |
| padding: 2px 10px; | |
| font-size: 0.78rem; | |
| font-weight: 700; | |
| letter-spacing: 0.04em; | |
| } | |
| .badge-worse { | |
| display: inline-block; | |
| background: #fee2e2; | |
| color: #991b1b; | |
| border: 1px solid #fca5a5; | |
| border-radius: 6px; | |
| padding: 2px 10px; | |
| font-size: 0.78rem; | |
| font-weight: 700; | |
| letter-spacing: 0.04em; | |
| } | |
| .badge-equal { | |
| display: inline-block; | |
| background: #f1f5f9; | |
| color: #64748b; | |
| border: 1px solid #cbd5e1; | |
| border-radius: 6px; | |
| padding: 2px 10px; | |
| font-size: 0.78rem; | |
| font-weight: 700; | |
| letter-spacing: 0.04em; | |
| } | |
| /* Print styles — hide Streamlit chrome when printing */ | |
| @media print { | |
| section[data-testid="stSidebar"], | |
| header[data-testid="stHeader"], | |
| .stToolbar, | |
| div[data-testid="stDecoration"], | |
| footer { display: none !important; } | |
| .main .block-container { padding: 0 !important; } | |
| } | |
| </style> | |
| """, | |
| unsafe_allow_html=True, | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Optional dependencies | |
| # --------------------------------------------------------------------------- | |
| try: | |
| import nibabel as nib # type: ignore[import] | |
| _NIB = True | |
| except ImportError: | |
| _NIB = False | |
| try: | |
| import plotly.graph_objects as go # type: ignore[import] | |
| from plotly.subplots import make_subplots # type: ignore[import] | |
| _PLOTLY = True | |
| except ImportError: | |
| _PLOTLY = False | |
| try: | |
| from scipy.ndimage import distance_transform_edt, binary_erosion # type: ignore[import] | |
| _SCIPY = True | |
| except ImportError: | |
| _SCIPY = False | |
| from segroc.model.models import ModelRegistry, SegmentationModel | |
| from segroc.data.perturbation import PerturbConfig, PerturbEngine | |
| from segroc.utils.metrics import ( | |
| auc_robustness, | |
| bootstrap_ci, | |
| psnr, | |
| ssim_score, | |
| wm_metric_t, | |
| m_ddeg_t, | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Constants | |
| # --------------------------------------------------------------------------- | |
| ARTIFACTS_DIR = Path("artifacts") | |
| PRECOMPUTED_DATASETS = ["covidquex", "isic2018", "kvasir", "mosmedplus", "promise12"] | |
| DATASET_DISPLAY: dict[str, str] = { | |
| "covidquex": "COVID-QuEX", | |
| "isic2018": "ISIC 2018 — Skin Lesion", | |
| "kvasir": "Kvasir — Polyp Segmentation", | |
| "mosmedplus": "MosMedPlus — COVID-19 CT", | |
| "promise12": "PROMISE12 — Prostate MR", | |
| } | |
| PRECOMPUTED_ARCHS: dict[str, str] = { | |
| "unet": "UNet", | |
| "attention_unet": "Attention UNet", | |
| "unetr": "UNETR", | |
| } | |
| PRECOMPUTED_SEEDS: list[int] = [11, 22, 33] | |
| ARTIFACT_OPTIONS: dict[str, list[str]] = { | |
| "CT": ["ring", "streak"], | |
| "MR": ["ghosting", "bias_field", "motion", "spike"], | |
| "Generic": ["noise"], | |
| } | |
| ARTIFACT_LABELS: dict[str, str] = { | |
| "ring": "Ring / Streak (CT detector)", | |
| "streak": "Metal Streak (CT)", | |
| "ghosting": "Motion Ghosting (MR)", | |
| "bias_field": "B1 Bias Field (MR)", | |
| "motion": "Rigid Motion (MR k-space)", | |
| "spike": "K-space Spike / Herringbone (MR)", | |
| "noise": "Gaussian Noise (Generic)", | |
| } | |
| _METRIC_META: dict[str, dict] = { | |
| "dice": {"label": "Dice", "axis": "left", "colour": "#2196F3"}, | |
| "hd95": {"label": "HD95 (vox)", "axis": "right", "colour": "#F44336"}, | |
| "hd100": {"label": "HD100 (vox)", "axis": "right", "colour": "#FF9800"}, | |
| "asd": {"label": "ASD (vox)", "axis": "right", "colour": "#9C27B0"}, | |
| "rmse": {"label": "RMSE (image)", "axis": "right", "colour": "#795548"}, | |
| "psnr": {"label": "PSNR dB", "axis": "right", "colour": "#4CAF50"}, | |
| "ssim": {"label": "SSIM (image)", "axis": "left", "colour": "#009688"}, | |
| } | |
| # --------------------------------------------------------------------------- | |
| # Precomputed bundle helpers | |
| # --------------------------------------------------------------------------- | |
| def get_precomputed_bundle_path(dataset: str, arch: str, seed: int) -> Path: | |
| return ARTIFACTS_DIR / f"{dataset}_{arch}_seed{seed}_bundle.zip" | |
| def load_precomputed_bundle(dataset: str, arch: str, seed: int) -> RobustnessBundle: | |
| path = get_precomputed_bundle_path(dataset, arch, seed) | |
| if not path.exists(): | |
| raise FileNotFoundError( | |
| f"Bundle not found at {path}. " | |
| "Ensure all artifact ZIPs are committed to the repository." | |
| ) | |
| return bundle_from_zip_bytes(path.read_bytes()) | |
| # --------------------------------------------------------------------------- | |
| # Sharpening (negative-severity direction) | |
| # --------------------------------------------------------------------------- | |
| def apply_sharpening(image: torch.Tensor, strength: float) -> torch.Tensor: | |
| """Unsharp-mask sharpening (2-D and 3-D tensors).""" | |
| if strength == 0.0: | |
| return image.clone() | |
| sigma = 1.5 | |
| ksize = 5 | |
| pad = ksize // 2 | |
| coords = torch.arange(ksize, dtype=torch.float32, device=image.device) - ksize // 2 | |
| g1d = torch.exp(-(coords**2) / (2 * sigma**2)) | |
| g1d = g1d / g1d.sum() | |
| img = image.float() | |
| if img.dim() <= 4: | |
| g2d = (g1d[:, None] * g1d[None, :]).view(1, 1, ksize, ksize) | |
| squeeze_back = img.dim() == 3 | |
| if squeeze_back: | |
| img = img.unsqueeze(0) | |
| blurred = F.conv2d( | |
| img, | |
| g2d.expand(img.shape[1], 1, ksize, ksize), | |
| padding=pad, | |
| groups=img.shape[1], | |
| ) | |
| sharpened = (img + strength * (img - blurred)).clamp(0.0, 1.0) | |
| if squeeze_back: | |
| sharpened = sharpened.squeeze(0) | |
| else: | |
| g3d = (g1d[:, None, None] * g1d[None, :, None] * g1d[None, None, :]).view( | |
| 1, 1, ksize, ksize, ksize | |
| ) | |
| squeeze_back = img.dim() == 4 | |
| if squeeze_back: | |
| img = img.unsqueeze(0) | |
| blurred = F.conv3d( | |
| img, | |
| g3d.expand(img.shape[1], 1, ksize, ksize, ksize), | |
| padding=pad, | |
| groups=img.shape[1], | |
| ) | |
| sharpened = (img + strength * (img - blurred)).clamp(0.0, 1.0) | |
| if squeeze_back: | |
| sharpened = sharpened.squeeze(0) | |
| return sharpened | |
| # --------------------------------------------------------------------------- | |
| # Checkpoint parsing & model loading | |
| # --------------------------------------------------------------------------- | |
| def _strip_prefix(raw: dict) -> dict: | |
| state = { | |
| k.replace("_model.", "", 1): v | |
| for k, v in raw.items() | |
| if k.startswith("_model.") | |
| } | |
| return state if state else raw | |
| def _infer_spatial_dims(state: dict) -> int: | |
| for v in state.values(): | |
| if isinstance(v, torch.Tensor) and v.dim() == 5: | |
| return 3 | |
| for v in state.values(): | |
| if isinstance(v, torch.Tensor) and v.dim() == 4: | |
| return 2 | |
| return 3 | |
| def _infer_in_channels(state: dict) -> int | None: | |
| for v in state.values(): | |
| if isinstance(v, torch.Tensor) and v.dim() in (4, 5): | |
| return int(v.shape[1]) | |
| return None | |
| def _candidate_out_channels(state: dict) -> list[int]: | |
| candidates = { | |
| int(v.shape[0]) | |
| for v in state.values() | |
| if isinstance(v, torch.Tensor) and v.dim() in (4, 5) | |
| } | |
| return sorted(candidates) | |
| def _infer_out_channels(state: dict) -> int | None: | |
| candidates = _candidate_out_channels(state) | |
| if not candidates: | |
| return None | |
| ordered = sorted(candidates, key=lambda x: (x > 8, x)) | |
| return ordered[0] | |
| def parse_checkpoint(ckpt_bytes: bytes) -> tuple[int | None, int | None, int]: | |
| buf = io.BytesIO(ckpt_bytes) | |
| ckpt = torch.load(buf, map_location="cpu", weights_only=False) | |
| raw = ckpt.get("state_dict", ckpt) | |
| state = _strip_prefix(raw) | |
| spatial_dims = _infer_spatial_dims(state) | |
| return _infer_in_channels(state), _infer_out_channels(state), spatial_dims | |
| def load_model_from_bytes( | |
| ckpt_bytes: bytes, | |
| model_name: str, | |
| in_ch: int, | |
| out_ch: int, | |
| spatial_dims: int = 3, | |
| ) -> SegmentationModel: | |
| buf = io.BytesIO(ckpt_bytes) | |
| ckpt = torch.load(buf, map_location="cpu", weights_only=False) | |
| raw = ckpt.get("state_dict", ckpt) | |
| state = _strip_prefix(raw) | |
| tried: list[int] = [] | |
| candidate_out = [out_ch] + [ | |
| c for c in _candidate_out_channels(state) if c != out_ch | |
| ] | |
| for out_candidate in candidate_out: | |
| tried.append(out_candidate) | |
| try: | |
| model = ModelRegistry.build( | |
| model_name, | |
| in_channels=in_ch, | |
| out_channels=out_candidate, | |
| spatial_dims=spatial_dims, | |
| ) | |
| model._backbone.load_state_dict(state, strict=True) | |
| return model.eval() | |
| except RuntimeError: | |
| continue | |
| raise RuntimeError( | |
| f"Unable to load checkpoint for model '{model_name}'. Tried out_channels={tried}." | |
| ) | |
| def get_demo_model( | |
| model_name: str, in_ch: int, out_ch: int, spatial_dims: int = 3 | |
| ) -> SegmentationModel: | |
| return ModelRegistry.build( | |
| model_name, in_channels=in_ch, out_channels=out_ch, spatial_dims=spatial_dims | |
| ).eval() | |
| # --------------------------------------------------------------------------- | |
| # NIfTI helpers | |
| # --------------------------------------------------------------------------- | |
| def _nifti_bytes_to_array_and_affine( | |
| nii_bytes: bytes, filename: str | |
| ) -> tuple[np.ndarray, np.ndarray]: | |
| if not _NIB: | |
| st.error("nibabel is required. `pip install nibabel`") | |
| st.stop() | |
| suffix = ".nii.gz" if filename.endswith(".nii.gz") else ".nii" | |
| with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp: | |
| tmp.write(nii_bytes) | |
| tmp_path = tmp.name | |
| try: | |
| img = nib.load(tmp_path) | |
| arr = np.asarray(img.dataobj, dtype=np.float32) | |
| affine = img.affine.copy() | |
| finally: | |
| Path(tmp_path).unlink(missing_ok=True) | |
| return arr, affine | |
| def load_nifti_bytes( | |
| nii_bytes: bytes, filename: str = "volume.nii.gz" | |
| ) -> tuple[torch.Tensor, np.ndarray, np.ndarray]: | |
| arr, affine = _nifti_bytes_to_array_and_affine(nii_bytes, filename) | |
| lo, hi = float(arr.min()), float(arr.max()) | |
| if hi > lo: | |
| arr = (arr - lo) / (hi - lo) | |
| tensor = torch.from_numpy(arr).unsqueeze(0).unsqueeze(0) | |
| return tensor, arr, affine | |
| def load_label_bytes(nii_bytes: bytes, filename: str = "label.nii.gz") -> torch.Tensor: | |
| arr, _ = _nifti_bytes_to_array_and_affine(nii_bytes, filename) | |
| return torch.from_numpy(arr).long() | |
| def save_nifti_bytes(arr: np.ndarray, affine: np.ndarray | None = None) -> bytes: | |
| if not _NIB: | |
| return b"" | |
| if affine is None: | |
| affine = np.eye(4) | |
| img = nib.Nifti1Image(arr, affine) | |
| with tempfile.NamedTemporaryFile(suffix=".nii.gz", delete=False) as tmp: | |
| tmp_path = tmp.name | |
| try: | |
| nib.save(img, tmp_path) | |
| with open(tmp_path, "rb") as f: | |
| return f.read() | |
| finally: | |
| Path(tmp_path).unlink(missing_ok=True) | |
| # --------------------------------------------------------------------------- | |
| # NPZ helpers (MedSegBench format) | |
| # --------------------------------------------------------------------------- | |
| def npz_test_count(npz_bytes: bytes) -> int: | |
| data = np.load(io.BytesIO(npz_bytes)) | |
| return int(data["test_images"].shape[0]) | |
| def load_npz_sample( | |
| npz_bytes: bytes, | |
| sample_idx: int, | |
| ) -> tuple[torch.Tensor, np.ndarray, np.ndarray | None]: | |
| data = np.load(io.BytesIO(npz_bytes)) | |
| test_images = data["test_images"] | |
| test_labels = data.get("test_label", None) | |
| img = test_images[sample_idx].astype(np.float32) | |
| lo, hi = float(img.min()), float(img.max()) | |
| if hi > lo: | |
| img = (img - lo) / (hi - lo) | |
| else: | |
| img = img / 255.0 | |
| if img.ndim == 2: | |
| img_chw = img[np.newaxis] | |
| display = img[np.newaxis] | |
| elif img.ndim == 3 and img.shape[2] in (1, 3, 4): | |
| img_chw = img.transpose(2, 0, 1) | |
| display = img.mean(axis=2)[np.newaxis] | |
| else: | |
| raise ValueError(f"Unexpected NPZ image shape: {img.shape}") | |
| tensor = torch.from_numpy(img_chw).float().unsqueeze(0) | |
| label: np.ndarray | None = None | |
| if test_labels is not None and sample_idx < len(test_labels): | |
| lbl = test_labels[sample_idx].astype(np.int64) | |
| label = (lbl > 0).astype(np.int64) | |
| return tensor, display, label | |
| # --------------------------------------------------------------------------- | |
| # PNG helpers | |
| # --------------------------------------------------------------------------- | |
| def load_png_bytes( | |
| png_bytes: bytes, filename: str = "image.png" | |
| ) -> tuple[torch.Tensor, np.ndarray]: | |
| try: | |
| from PIL import Image # type: ignore[import] | |
| except ImportError: | |
| st.error("PIL/Pillow is required. `pip install Pillow`") | |
| st.stop() | |
| try: | |
| img_pil = Image.open(io.BytesIO(png_bytes)) | |
| if img_pil.mode == "RGBA": | |
| img_pil = img_pil.convert("RGB") | |
| elif img_pil.mode not in ("L", "RGB"): | |
| img_pil = img_pil.convert("RGB") | |
| arr = np.asarray(img_pil, dtype=np.float32) | |
| except Exception as e: | |
| st.error(f"Failed to load PNG {filename}: {e}") | |
| st.stop() | |
| lo, hi = float(arr.min()), float(arr.max()) | |
| if hi > lo: | |
| arr = (arr - lo) / (hi - lo) | |
| else: | |
| arr = arr / 255.0 | |
| if arr.ndim == 2: | |
| img_chw = arr[np.newaxis] | |
| display = arr[np.newaxis] | |
| elif arr.ndim == 3: | |
| img_chw = arr.transpose(2, 0, 1) | |
| display = arr.mean(axis=2)[np.newaxis] | |
| else: | |
| raise ValueError(f"Unexpected PNG shape: {arr.shape}") | |
| tensor = torch.from_numpy(img_chw).float().unsqueeze(0) | |
| return tensor, display | |
| def load_png_label_bytes(png_bytes: bytes, filename: str = "label.png") -> np.ndarray: | |
| try: | |
| from PIL import Image # type: ignore[import] | |
| except ImportError: | |
| st.error("PIL/Pillow is required. `pip install Pillow`") | |
| st.stop() | |
| try: | |
| img_pil = Image.open(io.BytesIO(png_bytes)) | |
| if img_pil.mode != "L": | |
| img_pil = img_pil.convert("L") | |
| arr = np.asarray(img_pil, dtype=np.float32) | |
| except Exception as e: | |
| st.error(f"Failed to load PNG label {filename}: {e}") | |
| st.stop() | |
| lo, hi = float(arr.min()), float(arr.max()) | |
| if hi > lo: | |
| arr = (arr - lo) / (hi - lo) | |
| else: | |
| arr = arr / 255.0 | |
| return (arr > 0.5).astype(np.int64) | |
| # --------------------------------------------------------------------------- | |
| # Inference | |
| # --------------------------------------------------------------------------- | |
| def run_inference(model: nn.Module, image: torch.Tensor) -> torch.Tensor: | |
| if image.dim() == 4: | |
| logits = model(image) | |
| else: | |
| try: | |
| from monai.inferers import sliding_window_inference | |
| logits = sliding_window_inference(image, (64, 64, 64), 2, model) | |
| except Exception: | |
| logits = model(image) | |
| return logits.argmax(dim=1).squeeze(0) | |
| # --------------------------------------------------------------------------- | |
| # Metrics | |
| # --------------------------------------------------------------------------- | |
| def _hd_at_percentile(pred_bin: np.ndarray, gt_bin: np.ndarray, pct: float) -> float: | |
| if not pred_bin.any() or not gt_bin.any(): | |
| return float("nan") | |
| dt_gt = distance_transform_edt(~gt_bin) | |
| dt_pred = distance_transform_edt(~pred_bin) | |
| surf_pred = pred_bin & ~binary_erosion(pred_bin) | |
| surf_gt = gt_bin & ~binary_erosion(gt_bin) | |
| d1 = dt_gt[surf_pred] if surf_pred.any() else np.array([0.0]) | |
| d2 = dt_pred[surf_gt] if surf_gt.any() else np.array([0.0]) | |
| return float(np.percentile(np.concatenate([d1, d2]), pct)) | |
| def _mean_surface_distance(pred_bin: np.ndarray, gt_bin: np.ndarray) -> float: | |
| if not pred_bin.any() or not gt_bin.any(): | |
| return float("nan") | |
| dt_gt = distance_transform_edt(~gt_bin) | |
| dt_pred = distance_transform_edt(~pred_bin) | |
| surf_pred = pred_bin & ~binary_erosion(pred_bin) | |
| surf_gt = gt_bin & ~binary_erosion(gt_bin) | |
| d1 = dt_gt[surf_pred].mean() if surf_pred.any() else 0.0 | |
| d2 = dt_pred[surf_gt].mean() if surf_gt.any() else 0.0 | |
| return float((d1 + d2) / 2.0) | |
| def compute_all_metrics( | |
| pred: torch.Tensor, gt: torch.Tensor, n_classes: int | |
| ) -> dict[str, float]: | |
| pred_np = pred.cpu().numpy() | |
| gt_np = gt.cpu().numpy() | |
| dice_vals, hd95_vals, hd100_vals, asd_vals = [], [], [], [] | |
| for c in range(1, n_classes): | |
| pred_c = pred_np == c | |
| gt_c = gt_np == c | |
| if not gt_c.any(): | |
| continue | |
| inter = float((pred_c & gt_c).sum()) | |
| denom = float(pred_c.sum() + gt_c.sum()) | |
| dice_vals.append(2.0 * inter / denom if denom > 0 else 0.0) | |
| if _SCIPY: | |
| hd95_vals.append(_hd_at_percentile(pred_c, gt_c, 95)) | |
| hd100_vals.append(_hd_at_percentile(pred_c, gt_c, 100)) | |
| asd_vals.append(_mean_surface_distance(pred_c, gt_c)) | |
| result: dict[str, float] = { | |
| "dice": float(np.nanmean(dice_vals)) if dice_vals else float("nan"), | |
| } | |
| if _SCIPY: | |
| result["hd95"] = float(np.nanmean(hd95_vals)) if hd95_vals else float("nan") | |
| result["hd100"] = float(np.nanmean(hd100_vals)) if hd100_vals else float("nan") | |
| result["asd"] = float(np.nanmean(asd_vals)) if asd_vals else float("nan") | |
| return result | |
| # --------------------------------------------------------------------------- | |
| # Severity schedule & perturbation dispatch | |
| # --------------------------------------------------------------------------- | |
| def build_severity_schedule( | |
| n_sharp: int, n_degrade: int, max_sharp: float, max_degrade: float | |
| ) -> list[float]: | |
| sharp = [-max_sharp * (n_sharp - i) / n_sharp for i in range(n_sharp)] | |
| degrade = [max_degrade * (i + 1) / n_degrade for i in range(n_degrade)] | |
| return sharp + [0.0] + degrade | |
| def get_perturbed_image( | |
| image: torch.Tensor, | |
| signed_severity: float, | |
| artifact: str, | |
| modality: str, | |
| seed: int, | |
| ) -> torch.Tensor: | |
| if signed_severity == 0.0: | |
| return image.clone() | |
| if signed_severity < 0: | |
| return apply_sharpening(image, abs(signed_severity)) | |
| cfg = PerturbConfig( | |
| modality=modality, | |
| artifact=artifact, | |
| n_levels=2, | |
| seed=seed, | |
| intensity_range=(signed_severity, signed_severity), | |
| ) | |
| engine = PerturbEngine(cfg) | |
| if image.dim() == 4: | |
| img_for_perturb = image.squeeze(0).unsqueeze(-3) | |
| profile = engine.generate_profile(img_for_perturb) | |
| _, perturbed = profile.samples[1] | |
| return perturbed.squeeze(-3).unsqueeze(0) | |
| else: | |
| profile = engine.generate_profile(image.squeeze(0)) | |
| _, perturbed = profile.samples[1] | |
| return perturbed.unsqueeze(0) | |
| # --------------------------------------------------------------------------- | |
| # 3-D viewer helpers | |
| # --------------------------------------------------------------------------- | |
| def _slice_at(vol_3d: np.ndarray, axis: int, idx: int) -> np.ndarray: | |
| if axis == 0: | |
| return vol_3d[idx] | |
| if axis == 1: | |
| return vol_3d[:, idx, :] | |
| return vol_3d[:, :, idx] | |
| def _three_planes(vol_3d: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]: | |
| D, H, W = vol_3d.shape | |
| return vol_3d[D // 2], vol_3d[:, H // 2, :], vol_3d[:, :, W // 2] | |
| def _heatmap_fig( | |
| img: np.ndarray, | |
| overlay: np.ndarray | None = None, | |
| gt: np.ndarray | None = None, | |
| height: int = 220, | |
| ) -> "go.Figure": | |
| fig = go.Figure() | |
| fig.add_trace(go.Heatmap(z=img, colorscale="gray", showscale=False, name="Image")) | |
| if overlay is not None and overlay.max() > 0: | |
| fig.add_trace( | |
| go.Heatmap( | |
| z=overlay.astype(float), | |
| colorscale=[[0, "rgba(0,0,0,0)"], [1, "rgba(255,80,80,0.55)"]], | |
| showscale=False, | |
| name="Prediction", | |
| ) | |
| ) | |
| if gt is not None and gt.max() > 0: | |
| fig.add_trace( | |
| go.Heatmap( | |
| z=gt.astype(float), | |
| colorscale=[[0, "rgba(0,0,0,0)"], [1, "rgba(80,210,80,0.45)"]], | |
| showscale=False, | |
| name="GT", | |
| ) | |
| ) | |
| fig.update_layout( | |
| height=height, | |
| margin=dict(l=0, r=0, t=0, b=0), | |
| xaxis=dict(visible=False), | |
| yaxis=dict(visible=False, scaleanchor="x"), | |
| ) | |
| return fig | |
| def _render_three_planes( | |
| vol_3d: np.ndarray, | |
| pred_3d: np.ndarray | None = None, | |
| gt_3d: np.ndarray | None = None, | |
| show_seg: bool = True, | |
| height: int = 220, | |
| key_prefix: str = "", | |
| ax_idx: int | None = None, | |
| cor_idx: int | None = None, | |
| sag_idx: int | None = None, | |
| ) -> None: | |
| D, H, W = vol_3d.shape | |
| if D == 1: | |
| img_slice = vol_3d[0] | |
| pred_slice = (pred_3d[0] if pred_3d is not None else None) if show_seg else None | |
| gt_slice = (gt_3d[0] if gt_3d is not None else None) if show_seg else None | |
| st.caption("Image (2-D)") | |
| fig = _heatmap_fig(img_slice, pred_slice, gt_slice, height=height) | |
| st.plotly_chart(fig, use_container_width=True, key=f"{key_prefix}_2D") | |
| return | |
| ax_i = ax_idx if ax_idx is not None else D // 2 | |
| cor_i = cor_idx if cor_idx is not None else H // 2 | |
| sag_i = sag_idx if sag_idx is not None else W // 2 | |
| ax_img = vol_3d[ax_i] | |
| cor_img = vol_3d[:, cor_i, :] | |
| sag_img = vol_3d[:, :, sag_i] | |
| ax_pred = cor_pred = sag_pred = None | |
| ax_gt = cor_gt = sag_gt = None | |
| if show_seg and pred_3d is not None: | |
| ax_pred, cor_pred, sag_pred = ( | |
| pred_3d[ax_i], | |
| pred_3d[:, cor_i, :], | |
| pred_3d[:, :, sag_i], | |
| ) | |
| if show_seg and gt_3d is not None: | |
| ax_gt, cor_gt, sag_gt = gt_3d[ax_i], gt_3d[:, cor_i, :], gt_3d[:, :, sag_i] | |
| c1, c2, c3 = st.columns(3) | |
| for col, img_slice, pred_slice, gt_slice, plane in [ | |
| (c1, ax_img, ax_pred, ax_gt, "Axial"), | |
| (c2, cor_img, cor_pred, cor_gt, "Coronal"), | |
| (c3, sag_img, sag_pred, sag_gt, "Sagittal"), | |
| ]: | |
| with col: | |
| st.caption(plane) | |
| fig = _heatmap_fig(img_slice, pred_slice, gt_slice, height=height) | |
| st.plotly_chart(fig, use_container_width=True, key=f"{key_prefix}_{plane}") | |
| # --------------------------------------------------------------------------- | |
| # Synthetic demo | |
| # --------------------------------------------------------------------------- | |
| def make_synthetic_volume() -> tuple[torch.Tensor, torch.Tensor]: | |
| D, H, W = 64, 64, 64 | |
| image = torch.rand(1, 1, D, H, W) * 0.25 | |
| z, y, x = torch.meshgrid( | |
| torch.linspace(-1, 1, D), | |
| torch.linspace(-1, 1, H), | |
| torch.linspace(-1, 1, W), | |
| indexing="ij", | |
| ) | |
| fg = ((z / 0.55) ** 2 + (y / 0.65) ** 2 + (x / 0.45) ** 2) < 1.0 | |
| image[0, 0][fg] += 0.55 | |
| return image.clamp(0.0, 1.0), fg.long() | |
| def get_medsegbench_dataset_names() -> tuple[str, ...]: | |
| return list_medsegbench_datasets() | |
| def infer_dataset_modality(dataset_name: str) -> str: | |
| raw = dataset_name.lower().replace("-", "").replace("_", "") | |
| task_map = { | |
| task.name.lower().replace("-", "").replace("_", ""): task.modality | |
| for task in list_medsegbench_quickrun_tasks() | |
| } | |
| if raw in task_map: | |
| modality = task_map[raw] | |
| return modality if modality in ARTIFACT_OPTIONS else "Generic" | |
| alias_map = { | |
| "promise12": "MR", | |
| "mosmedplus": "CT", | |
| "kvasir": "Generic", | |
| "isic2018": "Generic", | |
| "covidquex": "Generic", | |
| } | |
| for key, modality in alias_map.items(): | |
| if key in raw: | |
| return modality | |
| return "Generic" | |
| def _safe_float_mean(values: list[float]) -> float: | |
| valid = [float(v) for v in values if not np.isnan(v)] | |
| return float(np.mean(valid)) if valid else float("nan") | |
| def _safe_corr(a: list[float], b: list[float]) -> float: | |
| aa = np.asarray(a, dtype=float) | |
| bb = np.asarray(b, dtype=float) | |
| mask = (~np.isnan(aa)) & (~np.isnan(bb)) | |
| if int(mask.sum()) < 2: | |
| return float("nan") | |
| aa, bb = aa[mask], bb[mask] | |
| if float(np.std(aa)) == 0.0 or float(np.std(bb)) == 0.0: | |
| return float("nan") | |
| return float(np.corrcoef(aa, bb)[0, 1]) | |
| def _align_input_channels(image: torch.Tensor, expected_channels: int) -> torch.Tensor: | |
| if image.dim() not in (4, 5): | |
| raise ValueError( | |
| f"Expected 4-D/5-D image tensor, got shape {tuple(image.shape)}" | |
| ) | |
| current_channels = int(image.shape[1]) | |
| if current_channels == expected_channels: | |
| return image | |
| if expected_channels == 1: | |
| return image.mean(dim=1, keepdim=True) | |
| if current_channels == 1 and expected_channels > 1: | |
| return image.repeat(1, expected_channels, *([1] * (image.dim() - 2))) | |
| if current_channels > expected_channels: | |
| return image[:, :expected_channels, ...] | |
| reps = int(np.ceil(expected_channels / current_channels)) | |
| expanded = image.repeat(1, reps, *([1] * (image.dim() - 2))) | |
| return expanded[:, :expected_channels, ...] | |
| def _binned_metric_envelope( | |
| df: pd.DataFrame, x_key: str, y_key: str, bins: int | |
| ) -> pd.DataFrame: | |
| if x_key not in df.columns or y_key not in df.columns: | |
| return pd.DataFrame(columns=["x", "y_mean", "y_min", "y_max", "count"]) | |
| sub = df[[x_key, y_key]].replace([np.inf, -np.inf], np.nan).dropna() | |
| if sub.empty: | |
| return pd.DataFrame(columns=["x", "y_mean", "y_min", "y_max", "count"]) | |
| xvals = sub[x_key].to_numpy(dtype=float) | |
| yvals = sub[y_key].to_numpy(dtype=float) | |
| x_min, x_max = float(np.min(xvals)), float(np.max(xvals)) | |
| if x_max <= x_min: | |
| return pd.DataFrame( | |
| { | |
| "x": [x_min], | |
| "y_mean": [float(np.nanmean(yvals))], | |
| "y_min": [float(np.nanmin(yvals))], | |
| "y_max": [float(np.nanmax(yvals))], | |
| "count": [len(yvals)], | |
| } | |
| ) | |
| n_bins = max(4, min(int(bins), 60)) | |
| edges = np.linspace(x_min, x_max, n_bins + 1) | |
| bucket = pd.cut(sub[x_key], bins=edges, include_lowest=True) | |
| grouped = ( | |
| sub.groupby(bucket, observed=False)[y_key] | |
| .agg(["mean", "min", "max", "count"]) | |
| .reset_index(drop=False) | |
| ) | |
| grouped = grouped[grouped["count"] > 0].copy() | |
| if grouped.empty: | |
| return pd.DataFrame(columns=["x", "y_mean", "y_min", "y_max", "count"]) | |
| centers = [ | |
| float((interval.left + interval.right) / 2.0) | |
| for interval in grouped[bucket.name].tolist() | |
| ] | |
| out = pd.DataFrame( | |
| { | |
| "x": centers, | |
| "y_mean": grouped["mean"].astype(float).tolist(), | |
| "y_min": grouped["min"].astype(float).tolist(), | |
| "y_max": grouped["max"].astype(float).tolist(), | |
| "count": grouped["count"].astype(int).tolist(), | |
| } | |
| ) | |
| return out.sort_values("x") | |
| # --------------------------------------------------------------------------- | |
| # Metric tooltip definitions | |
| # --------------------------------------------------------------------------- | |
| _METRIC_HELP: dict[str, str] = { | |
| "Dice": ( | |
| "Sørensen–Dice coefficient — measures the overlap between the predicted " | |
| "and ground-truth segmentation masks. Computed as 2×|P∩G| / (|P|+|G|). " | |
| "Range [0, 1]; higher is better." | |
| ), | |
| "HD95 (vox)": ( | |
| "95th-percentile Hausdorff distance in voxels — measures the worst-case " | |
| "boundary error while discarding the top 5 % of outlier surface points. " | |
| "Lower is better." | |
| ), | |
| "wmDSCt": ( | |
| "Weighted-mean Dice at transform level (Boone et al. 2023, NeuroImage) — " | |
| "a severity-weighted average of Dice across all degradation levels. " | |
| "Levels closer to the maximum severity receive higher weight, so a model " | |
| "that stays robust at high corruption scores well. Range [0, 1]; " | |
| "higher is better." | |
| ), | |
| "mDDegt": ( | |
| "Mean Dice degradation (Boone et al. 2023, NeuroImage) — average absolute " | |
| "drop in Dice relative to clean (unperturbed) performance across all " | |
| "degradation severity levels. A value of 0 means the model is unaffected " | |
| "by corruption; larger values indicate stronger degradation. " | |
| "Lower is better." | |
| ), | |
| "R-AUC": ( | |
| "Robustness Area Under the Curve — normalised area under the " | |
| "Dice-vs-severity curve over the degradation side. Captures the " | |
| "model's overall ability to maintain segmentation accuracy as corruption " | |
| "increases. Equivalent to 1 − mDDegt when the curve is linear. " | |
| "Range [0, 1]; higher is better." | |
| ), | |
| "Clean Dice": ( | |
| "Dice coefficient on unperturbed (clean) test images — the baseline " | |
| "segmentation accuracy with no added corruption. Higher is better." | |
| ), | |
| "Mean Dice": ( | |
| "Average Dice across all degradation severity levels for this artifact " | |
| "type. Gives a sense of typical performance under corruption. " | |
| "Higher is better." | |
| ), | |
| "Min Dice": ( | |
| "Minimum Dice observed at the highest tested severity level for this " | |
| "artifact type. Represents worst-case performance. Higher is better." | |
| ), | |
| "Clean HD95": ( | |
| "95th-percentile Hausdorff distance (voxels) on unperturbed images — " | |
| "the baseline boundary accuracy. Lower is better." | |
| ), | |
| "Mean HD95 (↑=worse)": ( | |
| "Average HD95 across all degradation severity levels for this artifact " | |
| "type. Higher values indicate worse boundary accuracy under corruption. " | |
| "Lower is better." | |
| ), | |
| "mDDegt (↑=worse)": ( | |
| "Mean Dice degradation for this artifact type — average drop in Dice " | |
| "relative to clean performance. Higher means more degradation. " | |
| "Lower is better." | |
| ), | |
| "🎯 Clean Dice": ( | |
| "Safety gate: the model's Dice on clean (unperturbed) images must meet " | |
| "this threshold before deployment is considered. Higher is better." | |
| ), | |
| "🦠 Corrupted Dice": ( | |
| "Safety gate: the macro-averaged Dice across all corruption types and " | |
| "severity levels must not drop by more than max_relative_dice_drop " | |
| "fraction of clean Dice. E.g. with a 10% relative limit and clean Dice " | |
| "0.90, the corrupted Dice must stay ≥ 0.81; with clean Dice 0.70 the " | |
| "bar is ≥ 0.63. This keeps the criterion proportional to baseline " | |
| "accuracy so harder tasks are not penalised unfairly. Higher is better." | |
| ), | |
| "📏 HD95": ( | |
| "Safety gate: the 95th-percentile Hausdorff distance under corruption " | |
| "must remain below this limit. " | |
| "⚠️ Measured in voxels (pixel units) — no voxel-to-mm spacing " | |
| "calibration is applied, so this value is NOT in millimetres regardless " | |
| "of field name. Interpret relative to the image resolution of the " | |
| "dataset under evaluation. Lower is better." | |
| ), | |
| "✔️ Pass Rate": ( | |
| "Safety gate: fraction of corrupted (sample × severity) evaluation " | |
| "points where Dice ≥ pass_rate_retention_fraction × clean_dice. The " | |
| "per-sample threshold adapts to the model's baseline: a model with " | |
| "clean Dice 0.90 at 80% retention must achieve ≥ 0.72 per sample, " | |
| "while one with clean Dice 0.70 must achieve ≥ 0.56. Coupling the bar " | |
| "to clean performance avoids penalising inherently harder tasks. " | |
| "Higher is better." | |
| ), | |
| } | |
| # --------------------------------------------------------------------------- | |
| # Fresh safety assessment (used to override baked-in bundle values) | |
| # --------------------------------------------------------------------------- | |
| def _fresh_safety( | |
| card: object, | |
| bundle_df: pd.DataFrame, | |
| ) -> tuple[SafetyAssessment, SafetyThresholds, float, dict[str, float]]: | |
| """Recompute safety gates and robustness score from current defaults. | |
| Returns (assessment, thresholds, overall_score, component_scores). | |
| overall_score is a continuous value in [0, 1] (0 = brittle, 1 = robust). | |
| component_scores keys: clean_dice, corrupted_dice, hd95, pass_rate, | |
| wm_dsc, m_ddeg, r_auc (only keys with available data are present). | |
| """ | |
| thresholds = SafetyThresholds() | |
| corrupt = bundle_df[bundle_df["direction"] == "degrade"]["dice"].dropna() | |
| # Per-sample pass threshold: each corrupted sample must retain at least | |
| # pass_rate_retention_fraction of the clean Dice score. | |
| pass_rate_threshold = ( | |
| float(card.clean_metrics.dice) * thresholds.pass_rate_retention_fraction | |
| ) | |
| pass_rate = ( | |
| float((corrupt.astype(float) >= pass_rate_threshold).mean()) | |
| if len(corrupt) > 0 | |
| else 0.0 | |
| ) | |
| assessment = assess_safety_gates( | |
| clean_dice=card.clean_metrics.dice, | |
| corrupted_mean_dice=card.corrupted_metrics_macro.dice, | |
| hd95_mm=card.corrupted_metrics_macro.hd95, | |
| pass_rate=pass_rate, | |
| thresholds=thresholds, | |
| ) | |
| # Per-component scores normalised to [0, 1] (1 = ideal). | |
| _HD95_WORST = 50.0 # voxels — anything above this is treated as maximally bad | |
| _DDEG_WORST = 1.0 # full Dice degradation | |
| components: dict[str, float] = {} | |
| components["clean_dice"] = float(np.clip(card.clean_metrics.dice, 0, 1)) | |
| components["corrupted_dice"] = float( | |
| np.clip(card.corrupted_metrics_macro.dice, 0, 1) | |
| ) | |
| components["pass_rate"] = float(np.clip(pass_rate, 0, 1)) | |
| hd95 = card.corrupted_metrics_macro.hd95 | |
| if hd95 is not None and not np.isnan(float(hd95)): | |
| components["hd95"] = float(np.clip(1 - float(hd95) / _HD95_WORST, 0, 1)) | |
| wm_dsc = card.corrupted_metrics_macro.wm_dsc | |
| if wm_dsc is not None and not np.isnan(float(wm_dsc)): | |
| components["wm_dsc"] = float(np.clip(float(wm_dsc), 0, 1)) | |
| m_ddeg = card.corrupted_metrics_macro.m_ddeg | |
| if m_ddeg is not None and not np.isnan(float(m_ddeg)): | |
| components["m_ddeg"] = float(np.clip(1 - float(m_ddeg) / _DDEG_WORST, 0, 1)) | |
| r_auc = card.corrupted_metrics_macro.r_auc | |
| if r_auc is not None and not np.isnan(float(r_auc)): | |
| components["r_auc"] = float(np.clip(float(r_auc), 0, 1)) | |
| # Weighted average; renormalise if some components are absent. | |
| _BASE_WEIGHTS: dict[str, float] = { | |
| "clean_dice": 0.20, | |
| "corrupted_dice": 0.25, | |
| "hd95": 0.15, | |
| "pass_rate": 0.20, | |
| "wm_dsc": 0.10, | |
| "m_ddeg": 0.05, | |
| "r_auc": 0.05, | |
| } | |
| total_w = sum(_BASE_WEIGHTS[k] for k in components if k in _BASE_WEIGHTS) | |
| if total_w > 0: | |
| overall = sum( | |
| components[k] * _BASE_WEIGHTS[k] / total_w | |
| for k in components | |
| if k in _BASE_WEIGHTS | |
| ) | |
| else: | |
| overall = 0.0 | |
| return assessment, thresholds, float(np.clip(overall, 0, 1)), components | |
| # --------------------------------------------------------------------------- | |
| # Score colouring helpers | |
| # --------------------------------------------------------------------------- | |
| def _score_color(s: float) -> tuple[str, str]: | |
| """Return (fg_hex, bg_hex) for a robustness score in [0, 1].""" | |
| if s >= 0.75: | |
| return "#14532d", "#dcfce7" # dark green / light green | |
| elif s >= 0.50: | |
| return "#854d0e", "#fef9c3" # dark amber / light yellow | |
| elif s >= 0.25: | |
| return "#9a3412", "#ffedd5" # dark orange / light orange | |
| else: | |
| return "#7f1d1d", "#fee2e2" # dark red / light red | |
| def _score_bar_html(s: float, height: int = 8) -> str: | |
| """Return an HTML progress bar coloured by score.""" | |
| fg, _ = _score_color(s) | |
| pct = f"{s * 100:.1f}" | |
| r = height // 2 | |
| return ( | |
| f'<div style="height:{height}px;background:#e5e7eb;border-radius:{r}px;overflow:hidden;">' | |
| f'<div style="height:{height}px;width:{pct}%;background:{fg};border-radius:{r}px;"></div>' | |
| f"</div>" | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Robustness card renderer | |
| # --------------------------------------------------------------------------- | |
| def _fmt_val(v: object) -> str: | |
| if v is None: | |
| return "—" | |
| try: | |
| f = float(v) # type: ignore[arg-type] | |
| return "—" if np.isnan(f) else f"{f:.4f}" | |
| except (TypeError, ValueError): | |
| return str(v) | |
| def _render_robustness_card_visual( | |
| card: object, | |
| assess_override: SafetyAssessment | None = None, | |
| thresholds_override: SafetyThresholds | None = None, | |
| robustness_score: float | None = None, | |
| component_scores: dict[str, float] | None = None, | |
| bundle_df: pd.DataFrame | None = None, | |
| ) -> None: | |
| """Render a visual robustness card (full standalone). | |
| Pass ``assess_override`` / ``thresholds_override`` to display a freshly | |
| computed safety assessment instead of the one baked into the bundle JSON. | |
| Pass ``robustness_score`` and ``component_scores`` (from ``_fresh_safety``) | |
| to show continuous scores instead of binary pass/fail. | |
| """ | |
| assess = assess_override if assess_override is not None else card.safety_assessment | |
| thresholds = ( | |
| thresholds_override | |
| if thresholds_override is not None | |
| else card.safety_thresholds | |
| ) | |
| score = robustness_score if robustness_score is not None else 0.0 | |
| comps = component_scores or {} | |
| score_fg, score_bg = _score_color(score) | |
| ts = card.generated_at_utc[:19].replace("T", " ") if card.generated_at_utc else "—" | |
| st.markdown( | |
| f"""<div style="background:{score_bg};border-left:6px solid {score_fg}; | |
| border-radius:8px;padding:16px 20px;margin-bottom:18px;"> | |
| <div style="display:flex;align-items:center;gap:18px;"> | |
| <div style="text-align:center;min-width:72px;"> | |
| <div style="font-size:2rem;font-weight:800;color:{score_fg}; | |
| line-height:1;">{score:.2f}</div> | |
| <div style="font-size:0.7rem;color:{score_fg};opacity:0.85; | |
| letter-spacing:0.03em;">ROBUSTNESS</div> | |
| </div> | |
| <div style="flex:1;"> | |
| {_score_bar_html(score, height=10)} | |
| <div style="color:#374151;font-size:0.78rem;margin-top:5px;"> | |
| 0 = brittle · 1 = no degradation | |
| </div> | |
| </div> | |
| </div> | |
| <div style="color:#6b7280;font-size:0.82rem;margin-top:10px;"> | |
| Model: <b>{card.model_name}</b> | | |
| Family: {card.model_family} | | |
| Generated (UTC): {ts} | | |
| Card v{card.card_version} | |
| </div> | |
| </div>""", | |
| unsafe_allow_html=True, | |
| ) | |
| st.markdown("##### Evaluation Scope") | |
| scope_df = pd.DataFrame( | |
| [ | |
| { | |
| "Dataset": card.dataset_scope.get("dataset", "—"), | |
| "Source": card.dataset_scope.get("source", "—"), | |
| "Split": card.dataset_scope.get("split", "—"), | |
| "N Samples": str(card.dataset_scope.get("n_samples", "—")), | |
| "Modality": card.perturbation_scope.get("modality", "—"), | |
| "Artifacts": ", ".join( | |
| str(a) for a in card.perturbation_scope.get("artifacts", []) | |
| ), | |
| "Severity Levels": str(card.perturbation_scope.get("n_levels", "—")), | |
| "Seed": str(card.perturbation_scope.get("seed", "—")), | |
| } | |
| ] | |
| ) | |
| st.dataframe(scope_df, hide_index=True, use_container_width=True) | |
| st.markdown("##### Performance Metrics") | |
| metrics_df = pd.DataFrame( | |
| [ | |
| { | |
| "Regime": "Clean", | |
| "Dice": _fmt_val(card.clean_metrics.dice), | |
| "HD95 (vox)": _fmt_val(card.clean_metrics.hd95), | |
| "wmDSCt": "—", | |
| "mDDegt": "—", | |
| "R-AUC": "—", | |
| }, | |
| { | |
| "Regime": "Corrupted (macro)", | |
| "Dice": _fmt_val(card.corrupted_metrics_macro.dice), | |
| "HD95 (vox)": _fmt_val(card.corrupted_metrics_macro.hd95), | |
| "wmDSCt": _fmt_val(card.corrupted_metrics_macro.wm_dsc), | |
| "mDDegt": _fmt_val(card.corrupted_metrics_macro.m_ddeg), | |
| "R-AUC": _fmt_val(card.corrupted_metrics_macro.r_auc), | |
| }, | |
| ] | |
| ) | |
| st.dataframe( | |
| metrics_df, | |
| hide_index=True, | |
| use_container_width=True, | |
| column_config={ | |
| "Dice": st.column_config.TextColumn("Dice", help=_METRIC_HELP["Dice"]), | |
| "HD95 (vox)": st.column_config.TextColumn( | |
| "HD95 (vox)", help=_METRIC_HELP["HD95 (vox)"] | |
| ), | |
| "wmDSCt": st.column_config.TextColumn( | |
| "wmDSCt", help=_METRIC_HELP["wmDSCt"] | |
| ), | |
| "mDDegt": st.column_config.TextColumn( | |
| "mDDegt", help=_METRIC_HELP["mDDegt"] | |
| ), | |
| "R-AUC": st.column_config.TextColumn("R-AUC", help=_METRIC_HELP["R-AUC"]), | |
| }, | |
| ) | |
| if card.artifact_breakdown: | |
| st.markdown("##### Artifact Breakdown") | |
| art_rows = [] | |
| for art, vals in sorted(card.artifact_breakdown.items()): | |
| art_rows.append( | |
| { | |
| "Artifact": art, | |
| "Clean Dice": _fmt_val(vals.get("clean_dice")), | |
| "Mean Dice": _fmt_val(vals.get("mean_dice")), | |
| "Min Dice": _fmt_val(vals.get("min_dice")), | |
| "wmDSCt": _fmt_val(vals.get("wm_dsc")), | |
| "mDDegt (↑=worse)": _fmt_val(vals.get("m_ddeg")), | |
| "R-AUC": _fmt_val(vals.get("r_auc")), | |
| "Clean HD95": _fmt_val(vals.get("clean_hd95")), | |
| "Mean HD95 (↑=worse)": _fmt_val(vals.get("mean_hd95")), | |
| } | |
| ) | |
| st.dataframe( | |
| pd.DataFrame(art_rows), | |
| hide_index=True, | |
| use_container_width=True, | |
| column_config={ | |
| "Clean Dice": st.column_config.TextColumn( | |
| "Clean Dice", help=_METRIC_HELP["Clean Dice"] | |
| ), | |
| "Mean Dice": st.column_config.TextColumn( | |
| "Mean Dice", help=_METRIC_HELP["Mean Dice"] | |
| ), | |
| "Min Dice": st.column_config.TextColumn( | |
| "Min Dice", help=_METRIC_HELP["Min Dice"] | |
| ), | |
| "wmDSCt": st.column_config.TextColumn( | |
| "wmDSCt", help=_METRIC_HELP["wmDSCt"] | |
| ), | |
| "mDDegt (↑=worse)": st.column_config.TextColumn( | |
| "mDDegt (↑=worse)", help=_METRIC_HELP["mDDegt (↑=worse)"] | |
| ), | |
| "R-AUC": st.column_config.TextColumn( | |
| "R-AUC", help=_METRIC_HELP["R-AUC"] | |
| ), | |
| "Clean HD95": st.column_config.TextColumn( | |
| "Clean HD95", help=_METRIC_HELP["Clean HD95"] | |
| ), | |
| "Mean HD95 (↑=worse)": st.column_config.TextColumn( | |
| "Mean HD95 (↑=worse)", help=_METRIC_HELP["Mean HD95 (↑=worse)"] | |
| ), | |
| }, | |
| ) | |
| st.markdown("##### Component Scores") | |
| # Compute per-sample pass threshold coupled to clean Dice performance. | |
| _prt = float(card.clean_metrics.dice) * thresholds.pass_rate_retention_fraction | |
| # Each gate: (label, comp_key, threshold_str, raw_value_str, help_key) | |
| safety_gates = [ | |
| ( | |
| "🎯 Clean Dice", | |
| "clean_dice", | |
| f"≥ {thresholds.min_clean_dice:.2f}", | |
| _fmt_val(card.clean_metrics.dice), | |
| "🎯 Clean Dice", | |
| ), | |
| ( | |
| "🦠 Corrupted Dice", | |
| "corrupted_dice", | |
| f"rel. drop ≤ {thresholds.max_relative_dice_drop:.0%}", | |
| _fmt_val(card.corrupted_metrics_macro.dice), | |
| "🦠 Corrupted Dice", | |
| ), | |
| ( | |
| "📏 HD95", | |
| "hd95", | |
| f"≤ {thresholds.max_hd95_mm:.1f} vox*", | |
| _fmt_val(card.corrupted_metrics_macro.hd95), | |
| "📏 HD95", | |
| ), | |
| ( | |
| "✔️ Pass Rate", | |
| "pass_rate", | |
| f"≥ {thresholds.min_pass_rate:.0%} (Dice≥{_prt:.2f})", | |
| "see eval", | |
| "✔️ Pass Rate", | |
| ), | |
| ] | |
| gate_cols = st.columns(len(safety_gates)) | |
| for col, (gate_name, comp_key, threshold, value, help_key) in zip( | |
| gate_cols, safety_gates | |
| ): | |
| with col: | |
| comp_s = comps.get(comp_key, 0.0) | |
| gc, gbg = _score_color(comp_s) | |
| tooltip = _METRIC_HELP.get(help_key, "") | |
| bar = _score_bar_html(comp_s, height=6) | |
| st.markdown( | |
| f"""<div title="{tooltip}" | |
| style="background:{gbg};border:1px solid {gc}44;border-radius:10px; | |
| padding:12px 10px;text-align:center;margin-bottom:4px;cursor:help;"> | |
| <div style="font-weight:600;color:#1f2937;font-size:0.85rem; | |
| margin-bottom:6px;">{gate_name}</div> | |
| <div style="font-size:1.55rem;font-weight:800;color:{gc}; | |
| line-height:1;">{comp_s:.2f}</div> | |
| <div style="margin:6px 4px 4px;">{bar}</div> | |
| <div style="color:#9ca3af;font-size:0.74rem;margin-top:4px;"> | |
| Threshold: {threshold} | |
| </div> | |
| <div style="color:#374151;font-size:0.79rem;font-weight:500;"> | |
| Value: {value} | |
| </div> | |
| </div>""", | |
| unsafe_allow_html=True, | |
| ) | |
| # ------------------------------------------------------------------ | |
| # Corrupted Dice distribution | |
| # ------------------------------------------------------------------ | |
| if ( | |
| bundle_df is not None | |
| and "direction" in bundle_df.columns | |
| and "dice" in bundle_df.columns | |
| ): | |
| degrade_dice = ( | |
| bundle_df[bundle_df["direction"] == "degrade"]["dice"] | |
| .dropna() | |
| .astype(float) | |
| .tolist() | |
| ) | |
| if degrade_dice: | |
| n_pts = len(degrade_dice) | |
| n_pass = int(sum(1 for d in degrade_dice if d >= _prt)) | |
| with st.expander( | |
| f"Corrupted Dice distribution — {n_pass}/{n_pts} points pass " | |
| f"(Dice ≥ {_prt:.3f})", | |
| expanded=False, | |
| ): | |
| st.caption( | |
| f"Distribution of Dice scores across all {n_pts} corrupted " | |
| f"(artifact × severity) evaluation points. " | |
| f"Dashed red line = per-sample pass threshold " | |
| f"({thresholds.pass_rate_retention_fraction:.0%} × clean Dice " | |
| f"{card.clean_metrics.dice:.3f} = {_prt:.3f}). " | |
| f"⚠️ This is a macro aggregate over severity levels — " | |
| f"inspect the Per-Perturbation tab to see worst-case severities." | |
| ) | |
| if _PLOTLY: | |
| _dist_fig = go.Figure() | |
| _dist_fig.add_trace( | |
| go.Histogram( | |
| x=degrade_dice, | |
| nbinsx=20, | |
| name="Corrupted Dice", | |
| marker_color="#3b82f6", | |
| opacity=0.75, | |
| ) | |
| ) | |
| _dist_fig.add_vline( | |
| x=_prt, | |
| line_dash="dash", | |
| line_color="#dc2626", | |
| annotation_text=f"threshold {_prt:.3f}", | |
| annotation_position="top left", | |
| ) | |
| _dist_fig.update_layout( | |
| height=220, | |
| margin=dict(l=0, r=0, t=20, b=0), | |
| xaxis_title="Dice", | |
| yaxis_title="Count", | |
| showlegend=False, | |
| ) | |
| st.plotly_chart(_dist_fig, use_container_width=True) | |
| else: | |
| st.bar_chart( | |
| pd.Series(degrade_dice, name="Corrupted Dice") | |
| .value_counts(bins=15, sort=False) | |
| .sort_index() | |
| ) | |
| st.caption(f"Per-sample pass threshold (not shown): {_prt:.3f}") | |
| if card.quality_linkage or card.governance_notes: | |
| with st.expander("Quality Linkage & Governance Notes"): | |
| if card.quality_linkage: | |
| st.markdown("**Quality-Metric Linkage**") | |
| for k, v in sorted(card.quality_linkage.items()): | |
| st.markdown(f"- **{k}**: {_fmt_val(v)}") | |
| if card.governance_notes: | |
| st.markdown("**Governance Notes**") | |
| for note in card.governance_notes: | |
| st.markdown(f"- {note}") | |
| # --------------------------------------------------------------------------- | |
| # Side-by-side comparison renderer | |
| # --------------------------------------------------------------------------- | |
| def _indicator_html( | |
| val_a: float, val_b: float, higher_is_better: bool | |
| ) -> tuple[str, str]: | |
| """Return (badge_html_a, badge_html_b) based on which value is better.""" | |
| if np.isnan(val_a) or np.isnan(val_b): | |
| return ( | |
| '<span class="badge-equal">N/A</span>', | |
| '<span class="badge-equal">N/A</span>', | |
| ) | |
| eps = 1e-6 | |
| if higher_is_better: | |
| a_better = val_a > val_b + eps | |
| b_better = val_b > val_a + eps | |
| else: | |
| a_better = val_a < val_b - eps | |
| b_better = val_b < val_a - eps | |
| if a_better: | |
| return ( | |
| '<span class="badge-better">▲ BETTER</span>', | |
| '<span class="badge-worse">▼ WORSE</span>', | |
| ) | |
| if b_better: | |
| return ( | |
| '<span class="badge-worse">▼ WORSE</span>', | |
| '<span class="badge-better">▲ BETTER</span>', | |
| ) | |
| return ( | |
| '<span class="badge-equal">= EQUAL</span>', | |
| '<span class="badge-equal">= EQUAL</span>', | |
| ) | |
| def _render_comparison_summary( | |
| card_a: object, card_b: object, label_a: str, label_b: str | |
| ) -> None: | |
| """Render a compact metric-by-metric comparison table at the top of Mode 2.""" | |
| rows = [] | |
| comparisons = [ | |
| ("Clean Dice", card_a.clean_metrics.dice, card_b.clean_metrics.dice, True), | |
| ( | |
| "Corrupted Dice", | |
| card_a.corrupted_metrics_macro.dice, | |
| card_b.corrupted_metrics_macro.dice, | |
| True, | |
| ), | |
| ( | |
| "Corrupted HD95", | |
| card_a.corrupted_metrics_macro.hd95, | |
| card_b.corrupted_metrics_macro.hd95, | |
| False, | |
| ), | |
| ( | |
| "wmDSCt", | |
| card_a.corrupted_metrics_macro.wm_dsc, | |
| card_b.corrupted_metrics_macro.wm_dsc, | |
| True, | |
| ), | |
| ( | |
| "mDDegt", | |
| card_a.corrupted_metrics_macro.m_ddeg, | |
| card_b.corrupted_metrics_macro.m_ddeg, | |
| False, | |
| ), | |
| ( | |
| "R-AUC", | |
| card_a.corrupted_metrics_macro.r_auc, | |
| card_b.corrupted_metrics_macro.r_auc, | |
| True, | |
| ), | |
| ] | |
| a_wins = b_wins = 0 | |
| for metric, val_a, val_b, hib in comparisons: | |
| badge_a_html, badge_b_html = _indicator_html( | |
| float(val_a) if val_a is not None else float("nan"), | |
| float(val_b) if val_b is not None else float("nan"), | |
| hib, | |
| ) | |
| if "BETTER" in badge_a_html: | |
| a_wins += 1 | |
| if "BETTER" in badge_b_html: | |
| b_wins += 1 | |
| rows.append( | |
| { | |
| "Metric": metric, | |
| f"{label_a}": _fmt_val(val_a), | |
| "": badge_a_html, | |
| f"{label_b}": _fmt_val(val_b), | |
| " ": badge_b_html, | |
| "Higher = Better": "✓" if hib else "✗", | |
| } | |
| ) | |
| winner = label_a if a_wins > b_wins else (label_b if b_wins > a_wins else "Tie") | |
| winner_color = "#16a34a" if winner != "Tie" else "#64748b" | |
| st.markdown( | |
| f"""<div style="background:linear-gradient(135deg,#f0fdf4,#dcfce7); | |
| border:1px solid #86efac;border-radius:12px;padding:16px 20px;margin-bottom:16px;"> | |
| <div style="font-size:1.05rem;font-weight:700;color:{winner_color};"> | |
| Overall winner: {winner} | |
| | | |
| <span style="color:#0f172a;font-weight:400;font-size:0.9rem;"> | |
| {label_a}: <b>{a_wins}</b> wins | | |
| {label_b}: <b>{b_wins}</b> wins (out of {len(comparisons)} metrics) | |
| </span> | |
| </div> | |
| </div>""", | |
| unsafe_allow_html=True, | |
| ) | |
| st.markdown("**Metric-by-metric comparison**") | |
| for row in rows: | |
| c1, c2, c3, c4, c5 = st.columns([2.5, 1.2, 1.4, 1.2, 1.4]) | |
| with c1: | |
| st.markdown(f"**{row['Metric']}**") | |
| with c2: | |
| st.markdown(row[label_a]) | |
| with c3: | |
| st.markdown(row[""], unsafe_allow_html=True) | |
| with c4: | |
| st.markdown(row[label_b]) | |
| with c5: | |
| st.markdown(row[" "], unsafe_allow_html=True) | |
| st.divider() | |
| # --------------------------------------------------------------------------- | |
| # Aggregate + per-perturbation chart helpers (shared across Mode 1 & 2) | |
| # --------------------------------------------------------------------------- | |
| def _render_quality_vs_seg_chart( | |
| card_df: pd.DataFrame, | |
| artifacts: list[str], | |
| key_prefix: str = "", | |
| ) -> None: | |
| """Aggregate view: image quality metric (x) vs segmentation metric (y).""" | |
| quality_metric_options = { | |
| "ssim": "SSIM (image)", | |
| "psnr": "PSNR dB (image)", | |
| "rmse": "RMSE (image)", | |
| } | |
| seg_metric_options = { | |
| "dice": "Dice", | |
| "hd95": "HD95 (vox)", | |
| "hd100": "HD100 (vox)", | |
| "asd": "ASD (vox)", | |
| } | |
| quality_choices = [k for k in ["ssim", "psnr", "rmse"] if k in card_df.columns] | |
| available_seg = [ | |
| k for k in ["dice", "hd95", "hd100", "asd"] if k in card_df.columns | |
| ] | |
| if not quality_choices or not available_seg: | |
| st.info("No image quality or segmentation metrics in this bundle.") | |
| return | |
| col1, col2, col3 = st.columns([2, 2, 1]) | |
| with col1: | |
| quality_key = st.selectbox( | |
| "Image quality metric (x-axis)", | |
| options=quality_choices, | |
| format_func=lambda k: quality_metric_options.get(k, k.upper()), | |
| key=f"{key_prefix}_qk", | |
| ) | |
| with col2: | |
| default_seg = [ | |
| k for k in ["dice", "hd95"] if k in available_seg | |
| ] or available_seg[:1] | |
| seg_keys = st.multiselect( | |
| "Segmentation metrics", | |
| options=available_seg, | |
| default=default_seg, | |
| format_func=lambda k: seg_metric_options.get(k, k), | |
| key=f"{key_prefix}_sk", | |
| ) | |
| with col3: | |
| n_bins = st.slider("Bins", 6, 36, 18, key=f"{key_prefix}_bins") | |
| show_points = st.toggle("Raw points", value=True, key=f"{key_prefix}_pts") | |
| if not seg_keys: | |
| st.warning("Select at least one segmentation metric.") | |
| return | |
| if _PLOTLY: | |
| fig = go.Figure() | |
| colour_map = { | |
| "dice": "#0f766e", | |
| "hd95": "#dc2626", | |
| "hd100": "#b45309", | |
| "asd": "#1d4ed8", | |
| } | |
| for seg_key in seg_keys: | |
| envelope = _binned_metric_envelope( | |
| card_df, quality_key, seg_key, bins=n_bins | |
| ) | |
| if envelope.empty: | |
| continue | |
| base_col = colour_map.get(seg_key, "#334155") | |
| rgb = f"{int(base_col[1:3],16)},{int(base_col[3:5],16)},{int(base_col[5:7],16)}" | |
| fig.add_trace( | |
| go.Scatter( | |
| x=envelope["x"].tolist(), | |
| y=envelope["y_max"].tolist(), | |
| mode="lines", | |
| line=dict(width=0), | |
| showlegend=False, | |
| hoverinfo="skip", | |
| legendgroup=seg_key, | |
| ) | |
| ) | |
| fig.add_trace( | |
| go.Scatter( | |
| x=envelope["x"].tolist(), | |
| y=envelope["y_min"].tolist(), | |
| mode="lines", | |
| fill="tonexty", | |
| fillcolor=f"rgba({rgb},0.14)", | |
| line=dict(width=0), | |
| name=f"{seg_metric_options.get(seg_key, seg_key)} band", | |
| legendgroup=seg_key, | |
| showlegend=True, | |
| ) | |
| ) | |
| fig.add_trace( | |
| go.Scatter( | |
| x=envelope["x"].tolist(), | |
| y=envelope["y_mean"].tolist(), | |
| mode="lines+markers", | |
| line=dict(color=base_col, width=2.5), | |
| marker=dict(size=6), | |
| name=f"{seg_metric_options.get(seg_key, seg_key)} mean", | |
| legendgroup=seg_key, | |
| ) | |
| ) | |
| if show_points: | |
| for art in artifacts: | |
| sub = card_df[card_df["artifact"] == art] | |
| if sub.empty: | |
| continue | |
| fig.add_trace( | |
| go.Scatter( | |
| x=sub[quality_key].tolist(), | |
| y=sub[seg_key].tolist(), | |
| mode="markers", | |
| marker=dict(size=5, color=base_col, opacity=0.22), | |
| name=f"{art}", | |
| legendgroup=f"{seg_key}_pts", | |
| showlegend=False, | |
| hovertemplate=f"Artifact: {art}<br>{quality_key}: %{{x:.4f}}<br>{seg_key}: %{{y:.4f}}<extra></extra>", | |
| ) | |
| ) | |
| fig.update_layout( | |
| height=520, | |
| template="plotly_white", | |
| hovermode="x unified", | |
| legend=dict(orientation="h", yanchor="bottom", y=1.02, x=0), | |
| margin=dict(l=20, r=20, t=20, b=20), | |
| ) | |
| fig.update_xaxes( | |
| title_text=quality_metric_options.get(quality_key, quality_key) | |
| ) | |
| fig.update_yaxes(title_text="Segmentation metrics") | |
| st.plotly_chart(fig, use_container_width=True) | |
| else: | |
| st.line_chart(card_df.set_index(quality_key)[seg_keys]) | |
| st.markdown("#### Aggregated metrics table") | |
| st.dataframe(card_df, use_container_width=True, hide_index=True) | |
| def _render_per_perturbation_chart( | |
| card_df: pd.DataFrame, | |
| artifacts: list[str], | |
| key_prefix: str = "", | |
| ) -> None: | |
| """Per-perturbation severity curves.""" | |
| art_selected = st.selectbox( | |
| "Perturbation type", | |
| options=artifacts, | |
| format_func=lambda a: ARTIFACT_LABELS.get(a, a), | |
| key=f"{key_prefix}_art", | |
| ) | |
| sub = card_df[card_df["artifact"] == art_selected].sort_values("severity") | |
| seg_options = [k for k in ["dice", "hd95", "hd100", "asd"] if k in sub.columns] | |
| qual_options = [k for k in ["ssim", "psnr", "rmse"] if k in sub.columns] | |
| p1, p2 = st.columns(2) | |
| with p1: | |
| per_seg = st.multiselect( | |
| "Segmentation lines", | |
| options=seg_options, | |
| default=[k for k in ["dice", "hd95"] if k in seg_options] | |
| or seg_options[:1], | |
| format_func=lambda k: { | |
| "dice": "Dice", | |
| "hd95": "HD95 (vox)", | |
| "hd100": "HD100 (vox)", | |
| "asd": "ASD (vox)", | |
| }.get(k, k), | |
| key=f"{key_prefix}_seg", | |
| ) | |
| with p2: | |
| per_quality = st.multiselect( | |
| "Image quality lines", | |
| options=qual_options, | |
| default=[k for k in ["ssim", "psnr"] if k in qual_options] | |
| or qual_options[:1], | |
| format_func=lambda k: { | |
| "ssim": "SSIM", | |
| "psnr": "PSNR dB", | |
| "rmse": "RMSE", | |
| }.get(k, k), | |
| key=f"{key_prefix}_qual", | |
| ) | |
| if _PLOTLY: | |
| fig_per = make_subplots(specs=[[{"secondary_y": True}]]) | |
| seg_colours = { | |
| "dice": "#0f766e", | |
| "hd95": "#dc2626", | |
| "hd100": "#b45309", | |
| "asd": "#1d4ed8", | |
| } | |
| qual_colours = {"ssim": "#0891b2", "psnr": "#16a34a", "rmse": "#a16207"} | |
| for key in per_seg: | |
| fig_per.add_trace( | |
| go.Scatter( | |
| x=sub["severity"].tolist(), | |
| y=sub[key].tolist(), | |
| mode="lines+markers", | |
| name={ | |
| "dice": "Dice", | |
| "hd95": "HD95", | |
| "hd100": "HD100", | |
| "asd": "ASD", | |
| }.get(key, key), | |
| line=dict(color=seg_colours.get(key, "#334155"), width=2.5), | |
| ), | |
| secondary_y=False, | |
| ) | |
| for key in per_quality: | |
| fig_per.add_trace( | |
| go.Scatter( | |
| x=sub["severity"].tolist(), | |
| y=sub[key].tolist(), | |
| mode="lines+markers", | |
| name={"ssim": "SSIM", "psnr": "PSNR", "rmse": "RMSE"}.get(key, key), | |
| line=dict( | |
| color=qual_colours.get(key, "#64748b"), width=2, dash="dot" | |
| ), | |
| ), | |
| secondary_y=True, | |
| ) | |
| fig_per.add_vline(x=0.0, line_dash="dash", line_color="rgba(100,116,139,0.6)") | |
| fig_per.update_layout( | |
| height=520, | |
| template="plotly_white", | |
| hovermode="x unified", | |
| title=ARTIFACT_LABELS.get(art_selected, art_selected), | |
| legend=dict(orientation="h", yanchor="bottom", y=1.02, x=0), | |
| ) | |
| fig_per.update_xaxes(title_text="Signed severity (− sharpen / + degrade)") | |
| fig_per.update_yaxes(title_text="Segmentation metrics", secondary_y=False) | |
| fig_per.update_yaxes( | |
| title_text="Image quality (SSIM / PSNR / RMSE)", secondary_y=True | |
| ) | |
| st.plotly_chart(fig_per, use_container_width=True) | |
| else: | |
| selected = per_seg + per_quality | |
| if selected: | |
| st.line_chart(sub.set_index("severity")[selected]) | |
| st.markdown("#### Selected perturbation metrics") | |
| st.dataframe(sub, use_container_width=True, hide_index=True) | |
| # --------------------------------------------------------------------------- | |
| # Print-as-PDF helper | |
| # --------------------------------------------------------------------------- | |
| def _print_as_pdf_button(label: str = "Print / Save as PDF") -> None: | |
| """Inject a browser-native print button (browser saves as PDF).""" | |
| _stc.html( | |
| f""" | |
| <button onclick="window.print()" style=" | |
| background:#0ea5e9;color:#fff;border:none;border-radius:8px; | |
| padding:9px 20px;font-size:0.92rem;font-weight:600; | |
| cursor:pointer;letter-spacing:0.03em; | |
| box-shadow:0 1px 4px rgba(0,0,0,0.12); | |
| transition:background 0.15s; | |
| " | |
| onmouseover="this.style.background='#0284c7'" | |
| onmouseout="this.style.background='#0ea5e9'" | |
| >🖨 {label}</button> | |
| """, | |
| height=50, | |
| ) | |
| # =========================================================================== | |
| # SIDEBAR | |
| # =========================================================================== | |
| with st.sidebar: | |
| st.markdown( | |
| """<div style="text-align:center;padding:12px 0 4px 0;"> | |
| <span style="font-size:2rem;font-weight:900;color:#e2e8f0;letter-spacing:-1px;">segroc</span> | |
| <div style="font-size:0.78rem;color:#64748b;letter-spacing:0.08em;text-transform:uppercase;"> | |
| Segmentation Robustness Explorer | |
| </div> | |
| </div>""", | |
| unsafe_allow_html=True, | |
| ) | |
| st.divider() | |
| eval_mode = st.radio( | |
| "Evaluation Mode", | |
| ["MedSegBench Robustness Cards", "Compare Models", "Upload Your Model"], | |
| index=0, | |
| key="eval_mode_radio", | |
| ) | |
| st.divider() | |
| # ── Shared defaults ────────────────────────────────────────────────────── | |
| ckpt_bytes: bytes | None = None | |
| final_arch = "unet" | |
| final_in_ch = 1 | |
| final_out_ch = 2 | |
| final_spatial = 3 | |
| nii_file = npz_file = png_file = png_label_file = gt_file = None | |
| input_mode = "NIfTI (.nii / .nii.gz)" | |
| npz_sample_idx = 0 | |
| modality = "Generic" | |
| artifact = "noise" | |
| n_sharp = 8 | |
| n_degrade = 8 | |
| max_sharp = 2.0 | |
| max_degrade = 3.0 | |
| seed = 42 | |
| run_btn = False | |
| # ── Mode 1: MedSegBench Robustness Cards ───────────────────────────────── | |
| if eval_mode == "MedSegBench Robustness Cards": | |
| st.markdown( | |
| "<div style='color:#94a3b8;font-size:0.82rem;text-transform:uppercase;" | |
| "letter-spacing:0.06em;margin-bottom:8px;'>Select a Model Run</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| explore_dataset = st.selectbox( | |
| "Dataset", | |
| PRECOMPUTED_DATASETS, | |
| format_func=lambda d: DATASET_DISPLAY.get(d, d), | |
| key="exp_dataset", | |
| ) | |
| explore_arch = st.selectbox( | |
| "Architecture", | |
| list(PRECOMPUTED_ARCHS.keys()), | |
| format_func=lambda a: PRECOMPUTED_ARCHS.get(a, a), | |
| key="exp_arch", | |
| ) | |
| explore_seed = st.selectbox( | |
| "Training Seed", | |
| PRECOMPUTED_SEEDS, | |
| format_func=lambda s: f"Seed {s}", | |
| key="exp_seed", | |
| ) | |
| # ── Mode 2: Compare Models ─────────────────────────────────────────────── | |
| elif eval_mode == "Compare Models": | |
| st.markdown( | |
| "<div style='color:#94a3b8;font-size:0.82rem;text-transform:uppercase;" | |
| "letter-spacing:0.06em;margin-bottom:8px;'>Shared Dataset</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| cmp_dataset = st.selectbox( | |
| "Dataset", | |
| PRECOMPUTED_DATASETS, | |
| format_func=lambda d: DATASET_DISPLAY.get(d, d), | |
| key="cmp_dataset", | |
| help="Both models are evaluated on the same dataset — this ensures the comparison is fair.", | |
| ) | |
| # Expose cmp_a_dataset / cmp_b_dataset as aliases so downstream code is unchanged. | |
| cmp_a_dataset = cmp_dataset | |
| cmp_b_dataset = cmp_dataset | |
| st.divider() | |
| st.markdown( | |
| "<div style='color:#94a3b8;font-size:0.82rem;text-transform:uppercase;" | |
| "letter-spacing:0.06em;margin-bottom:8px;'>Model A</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| cmp_a_arch = st.selectbox( | |
| "Architecture A", | |
| list(PRECOMPUTED_ARCHS.keys()), | |
| format_func=lambda a: PRECOMPUTED_ARCHS.get(a, a), | |
| key="cmp_a_arch", | |
| ) | |
| cmp_a_seed = st.selectbox( | |
| "Seed A", | |
| PRECOMPUTED_SEEDS, | |
| format_func=lambda s: f"Seed {s}", | |
| key="cmp_a_seed", | |
| ) | |
| st.divider() | |
| st.markdown( | |
| "<div style='color:#94a3b8;font-size:0.82rem;text-transform:uppercase;" | |
| "letter-spacing:0.06em;margin-bottom:8px;'>Model B</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| cmp_b_arch = st.selectbox( | |
| "Architecture B", | |
| list(PRECOMPUTED_ARCHS.keys()), | |
| index=1, | |
| format_func=lambda a: PRECOMPUTED_ARCHS.get(a, a), | |
| key="cmp_b_arch", | |
| ) | |
| cmp_b_seed = st.selectbox( | |
| "Seed B", | |
| PRECOMPUTED_SEEDS, | |
| key="cmp_b_seed", | |
| ) | |
| # ── Mode 3: Upload Your Model ──────────────────────────────────────────── | |
| else: | |
| st.markdown( | |
| """<div class="segroc-disclaimer" style="margin:0 0 12px 0;"> | |
| ⚠️ <b>MONAI models only.</b> Only architectures from the | |
| MONAI framework are supported (UNet, Attention UNet, UNETR, | |
| SwinUNETR, VNet). The architecture is inferred from the .pt | |
| file structure; you must confirm the correct architecture | |
| below. | |
| </div>""", | |
| unsafe_allow_html=True, | |
| ) | |
| arch_list = ModelRegistry.list_models() | |
| monai_archs = [ | |
| m | |
| for m in arch_list | |
| if m in ["unet", "attention_unet", "unetr", "swin_unetr", "vnet"] | |
| ] or arch_list | |
| st.subheader("Model Checkpoint") | |
| ckpt_file = st.file_uploader( | |
| "MONAI checkpoint (.pt / .ckpt) — required", | |
| type=["ckpt", "pt", "pth"], | |
| key="upload_ckpt", | |
| ) | |
| if ckpt_file is not None: | |
| ckpt_bytes = ckpt_file.getvalue() | |
| inferred_in, inferred_out, inferred_spatial = parse_checkpoint(ckpt_bytes) | |
| st.info( | |
| f"**Inferred from checkpoint:** \n" | |
| f"• in_channels: {inferred_in} \n" | |
| f"• out_channels (best guess): {inferred_out} \n" | |
| f"• spatial_dims: {inferred_spatial}D" | |
| ) | |
| model_name = st.selectbox( | |
| "Confirm Architecture", | |
| monai_archs, | |
| help="Select the MONAI architecture that matches the checkpoint.", | |
| key="upload_arch_sel", | |
| ) | |
| final_arch = model_name | |
| final_in_ch = inferred_in or 1 | |
| final_out_ch = inferred_out or 2 | |
| final_spatial = inferred_spatial | |
| else: | |
| model_name = st.selectbox( | |
| "Architecture (demo — random weights)", | |
| monai_archs, | |
| key="upload_arch_sel", | |
| ) | |
| final_arch = model_name | |
| st.caption("No checkpoint — demo mode with random weights.") | |
| st.divider() | |
| st.subheader("Input Image") | |
| input_mode = st.radio( | |
| "Image format", | |
| ["NIfTI (.nii / .nii.gz)", "NPZ — MedSegBench", "PNG Image"], | |
| key="upload_input_mode", | |
| ) | |
| if input_mode == "NIfTI (.nii / .nii.gz)": | |
| nii_file = st.file_uploader( | |
| "Image (.nii / .nii.gz) — leave empty for synthetic demo", | |
| type=["nii", "gz"], | |
| key="upload_nii", | |
| ) | |
| gt_file = st.file_uploader( | |
| "Ground-truth label (optional)", | |
| type=["nii", "gz"], | |
| key="upload_gt", | |
| ) | |
| elif input_mode == "NPZ — MedSegBench": | |
| npz_file = st.file_uploader( | |
| "MedSegBench NPZ (e.g. promise12_128.npz)", | |
| type=["npz"], | |
| key="upload_npz", | |
| ) | |
| if npz_file is not None: | |
| _npz_bytes_peek = npz_file.getvalue() | |
| _n_test = npz_test_count(_npz_bytes_peek) | |
| st.caption(f"Test split: **{_n_test}** samples") | |
| npz_sample_idx = int( | |
| st.slider("Sample index", 0, max(0, _n_test - 1), 0) | |
| ) | |
| else: | |
| png_file = st.file_uploader("PNG image", type=["png"], key="upload_png") | |
| png_label_file = st.file_uploader( | |
| "PNG label mask (optional)", type=["png"], key="upload_png_lbl" | |
| ) | |
| modality = st.selectbox( | |
| "Modality", ["CT", "MR", "Generic"], index=1, key="upload_mod" | |
| ) | |
| st.divider() | |
| st.subheader("Perturbation") | |
| artifact_choices = ARTIFACT_OPTIONS[modality] | |
| artifact = st.selectbox( | |
| "Degradation artifact", | |
| artifact_choices, | |
| format_func=lambda k: ARTIFACT_LABELS[k], | |
| key="upload_artifact", | |
| ) | |
| st.caption("Negative severity = sharpening; positive = artifact degradation.") | |
| col_l, col_r = st.columns(2) | |
| with col_l: | |
| n_sharp = st.slider("Sharp levels", 1, 15, 8, key="upload_nsharp") | |
| with col_r: | |
| n_degrade = st.slider("Degrade levels", 1, 15, 8, key="upload_ndegrade") | |
| max_sharp = st.slider( | |
| "Max sharpening", 0.5, 5.0, 2.0, 0.25, key="upload_maxsharp" | |
| ) | |
| max_degrade = st.slider( | |
| "Max degradation", 0.5, 5.0, 3.0, 0.25, key="upload_maxdegrade" | |
| ) | |
| seed = st.number_input("Random seed", value=42, step=1, key="upload_seed") | |
| st.divider() | |
| run_btn = st.button( | |
| "Run Evaluation", | |
| type="primary", | |
| use_container_width=True, | |
| key="upload_run_btn", | |
| ) | |
| # =========================================================================== | |
| # MODE 1: EXPLORE PRECOMPUTED | |
| # =========================================================================== | |
| if eval_mode == "MedSegBench Robustness Cards": | |
| st.markdown( | |
| f"""<div class="segroc-hero"> | |
| <h2>MedSegBench Robustness Cards</h2> | |
| <p>A demonstration of pre-computed robustness evaluations across | |
| <b>5 MedSegBench datasets</b>, <b>3 architectures</b> | |
| (UNet · Attention UNet · UNETR), and <b>3 training seeds</b> | |
| — browse instantly, no compute required.</p> | |
| </div>""", | |
| unsafe_allow_html=True, | |
| ) | |
| try: | |
| bundle = load_precomputed_bundle(explore_dataset, explore_arch, explore_seed) | |
| except FileNotFoundError as _e: | |
| st.error(str(_e)) | |
| st.stop() | |
| except Exception as _e: | |
| st.error(f"Failed to load bundle: {_e}") | |
| st.stop() | |
| card = bundle.card | |
| card_df: pd.DataFrame = bundle_to_streamlit_eval_results(bundle)["df"] | |
| artifacts: list[str] = bundle.artifacts | |
| dataset_name: str = bundle.dataset | |
| modality_name: str = bundle.modality | |
| n_samples: int = bundle.n_samples | |
| k1, k2, k3, k4 = st.columns(4) | |
| with k1: | |
| st.metric("Clean Dice", f"{card.clean_metrics.dice:.4f}") | |
| with k2: | |
| st.metric("Corrupted Dice", f"{card.corrupted_metrics_macro.dice:.4f}") | |
| with k3: | |
| hd95_val = card.corrupted_metrics_macro.hd95 | |
| st.metric( | |
| "Corrupted HD95", | |
| ( | |
| "N/A" | |
| if hd95_val is None or np.isnan(float(hd95_val)) | |
| else f"{float(hd95_val):.4f}" | |
| ), | |
| ) | |
| _exp_assess, _exp_thresholds, _exp_score, _exp_comps = _fresh_safety(card, card_df) | |
| with k4: | |
| st.metric("Robustness Score", f"{_exp_score:.2f}") | |
| tab_card, tab_quality, tab_perturb, tab_dl = st.tabs( | |
| [ | |
| "Robustness Card", | |
| "Aggregate by Image Quality", | |
| "Per-Perturbation", | |
| "Downloads", | |
| ] | |
| ) | |
| with tab_card: | |
| _render_robustness_card_visual( | |
| card, | |
| _exp_assess, | |
| _exp_thresholds, | |
| _exp_score, | |
| _exp_comps, | |
| bundle_df=card_df, | |
| ) | |
| with tab_quality: | |
| st.markdown( | |
| "<div class='segroc-note'><b>Aggregate view:</b> each point is one " | |
| "perturbation type × severity level. The line shows the binned mean; " | |
| "the band captures variability across artifact types at similar " | |
| "image quality.</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| _render_quality_vs_seg_chart(card_df, artifacts, key_prefix="exp") | |
| with tab_perturb: | |
| st.markdown( | |
| "<div class='segroc-note'><b>Per-perturbation view:</b> x-axis is " | |
| "signed severity (negative = sharpening, positive = degradation). " | |
| "Use this to compare how different artifact families stress the " | |
| "model.</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| _render_per_perturbation_chart(card_df, artifacts, key_prefix="exp") | |
| with tab_dl: | |
| st.markdown("#### Download Robustness Card") | |
| col_dl1, col_dl2, col_dl3 = st.columns(3) | |
| _ev_results = bundle_to_streamlit_eval_results(bundle) | |
| with col_dl1: | |
| st.download_button( | |
| label="Download card JSON", | |
| data=_ev_results["card_json"], | |
| file_name=f"robustness_card_{dataset_name}_{explore_arch}_seed{explore_seed}.json", | |
| mime="application/json", | |
| use_container_width=True, | |
| ) | |
| with col_dl2: | |
| st.download_button( | |
| label="Download metrics CSV", | |
| data=card_df.to_csv(index=False).encode(), | |
| file_name=f"robustness_metrics_{dataset_name}_{explore_arch}_seed{explore_seed}.csv", | |
| mime="text/csv", | |
| use_container_width=True, | |
| ) | |
| with col_dl3: | |
| st.download_button( | |
| label="Download bundle ZIP", | |
| data=bundle_to_zip_bytes(bundle), | |
| file_name=f"robustness_bundle_{dataset_name}_{explore_arch}_seed{explore_seed}.zip", | |
| mime="application/zip", | |
| use_container_width=True, | |
| ) | |
| st.markdown("#### Print / Save as PDF") | |
| st.markdown( | |
| "<div class='segroc-note'>Click the button below to open your browser's " | |
| "print dialog. Select <b>Save as PDF</b> as the destination to export the " | |
| "current page (including the robustness card tab) as a PDF document.</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| _print_as_pdf_button("Print current page as PDF") | |
| st.stop() | |
| # =========================================================================== | |
| # MODE 2: COMPARE MODELS | |
| # =========================================================================== | |
| if eval_mode == "Compare Models": | |
| st.markdown( | |
| f"""<div class="segroc-hero"> | |
| <h2>Compare Models</h2> | |
| <p>Select two models evaluated on the same dataset | |
| to view their robustness cards side-by-side with | |
| per-metric win/loss indicators.</p> | |
| </div>""", | |
| unsafe_allow_html=True, | |
| ) | |
| try: | |
| bundle_a = load_precomputed_bundle(cmp_a_dataset, cmp_a_arch, cmp_a_seed) | |
| bundle_b = load_precomputed_bundle(cmp_b_dataset, cmp_b_arch, cmp_b_seed) | |
| except FileNotFoundError as _e: | |
| st.error(str(_e)) | |
| st.stop() | |
| except Exception as _e: | |
| st.error(f"Failed to load one or both bundles: {_e}") | |
| st.stop() | |
| _dataset_display_name = DATASET_DISPLAY.get(cmp_a_dataset, cmp_a_dataset) | |
| label_a = f"{PRECOMPUTED_ARCHS.get(cmp_a_arch, cmp_a_arch)} / Seed {cmp_a_seed}" | |
| label_b = f"{PRECOMPUTED_ARCHS.get(cmp_b_arch, cmp_b_arch)} / Seed {cmp_b_seed}" | |
| st.caption(f"Dataset: **{_dataset_display_name}**") | |
| _render_comparison_summary(bundle_a.card, bundle_b.card, "Model A", "Model B") | |
| _ev_a = bundle_to_streamlit_eval_results(bundle_a) | |
| _ev_b = bundle_to_streamlit_eval_results(bundle_b) | |
| tab_sidebyside, tab_curves, tab_dl_cmp = st.tabs( | |
| [ | |
| "Side-by-Side Cards", | |
| "Severity Curves", | |
| "Downloads", | |
| ] | |
| ) | |
| with tab_sidebyside: | |
| col_a, col_b = st.columns(2) | |
| with col_a: | |
| st.markdown( | |
| f"<div style='background:#f0fdf4;border:1px solid #86efac;" | |
| f"border-radius:8px;padding:10px 14px;margin-bottom:12px;'>" | |
| f"<b style='color:#166534;font-size:0.95rem;'>Model A</b><br>" | |
| f"<span style='color:#374151;font-size:0.85rem;'>{label_a}</span>" | |
| f"</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| _cmp_a_assess, _cmp_a_thresh, _cmp_a_score, _cmp_a_comps = _fresh_safety( | |
| bundle_a.card, _ev_a["df"] | |
| ) | |
| _render_robustness_card_visual( | |
| bundle_a.card, | |
| _cmp_a_assess, | |
| _cmp_a_thresh, | |
| _cmp_a_score, | |
| _cmp_a_comps, | |
| bundle_df=_ev_a["df"], | |
| ) | |
| with col_b: | |
| st.markdown( | |
| f"<div style='background:#eff6ff;border:1px solid #93c5fd;" | |
| f"border-radius:8px;padding:10px 14px;margin-bottom:12px;'>" | |
| f"<b style='color:#1d4ed8;font-size:0.95rem;'>Model B</b><br>" | |
| f"<span style='color:#374151;font-size:0.85rem;'>{label_b}</span>" | |
| f"</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| _cmp_b_assess, _cmp_b_thresh, _cmp_b_score, _cmp_b_comps = _fresh_safety( | |
| bundle_b.card, _ev_b["df"] | |
| ) | |
| _render_robustness_card_visual( | |
| bundle_b.card, | |
| _cmp_b_assess, | |
| _cmp_b_thresh, | |
| _cmp_b_score, | |
| _cmp_b_comps, | |
| bundle_df=_ev_b["df"], | |
| ) | |
| with tab_curves: | |
| st.markdown( | |
| "<div class='segroc-note'>Per-perturbation severity curves for both " | |
| "models. Select a dataset and perturbation type to compare how each " | |
| "model degrades under increasing corruption.</div>", | |
| unsafe_allow_html=True, | |
| ) | |
| _df_a: pd.DataFrame = _ev_a["df"] | |
| _df_b: pd.DataFrame = _ev_b["df"] | |
| _art_opts_a = sorted(_df_a["artifact"].astype(str).unique().tolist()) | |
| _art_opts_b = sorted(_df_b["artifact"].astype(str).unique().tolist()) | |
| _art_union = sorted(set(_art_opts_a) | set(_art_opts_b)) | |
| _sel_art = st.selectbox( | |
| "Perturbation type", | |
| _art_union, | |
| format_func=lambda a: ARTIFACT_LABELS.get(a, a), | |
| key="cmp_sel_art", | |
| ) | |
| _cmp_seg_choices = [ | |
| m | |
| for m in ["dice", "hd95", "hd100", "asd"] | |
| if m in _df_a.columns or m in _df_b.columns | |
| ] | |
| _cmp_metric = st.selectbox( | |
| "Segmentation metric", | |
| _cmp_seg_choices, | |
| format_func=lambda k: { | |
| "dice": "Dice", | |
| "hd95": "HD95", | |
| "hd100": "HD100", | |
| "asd": "ASD", | |
| }.get(k, k), | |
| key="cmp_metric", | |
| ) | |
| _sub_a = _df_a[_df_a["artifact"] == _sel_art].sort_values("severity") | |
| _sub_b = _df_b[_df_b["artifact"] == _sel_art].sort_values("severity") | |
| if _PLOTLY: | |
| _fig_cmp = go.Figure() | |
| if not _sub_a.empty and _cmp_metric in _sub_a.columns: | |
| _fig_cmp.add_trace( | |
| go.Scatter( | |
| x=_sub_a["severity"].tolist(), | |
| y=_sub_a[_cmp_metric].tolist(), | |
| mode="lines+markers", | |
| name=f"A — {PRECOMPUTED_ARCHS.get(cmp_a_arch, cmp_a_arch)} / Seed {cmp_a_seed}", | |
| line=dict(color="#0f766e", width=2.5), | |
| ) | |
| ) | |
| if not _sub_b.empty and _cmp_metric in _sub_b.columns: | |
| _fig_cmp.add_trace( | |
| go.Scatter( | |
| x=_sub_b["severity"].tolist(), | |
| y=_sub_b[_cmp_metric].tolist(), | |
| mode="lines+markers", | |
| name=f"B — {PRECOMPUTED_ARCHS.get(cmp_b_arch, cmp_b_arch)} / Seed {cmp_b_seed}", | |
| line=dict(color="#1d4ed8", width=2.5), | |
| ) | |
| ) | |
| _fig_cmp.add_vline( | |
| x=0.0, line_dash="dash", line_color="rgba(100,116,139,0.6)" | |
| ) | |
| _fig_cmp.update_layout( | |
| height=480, | |
| template="plotly_white", | |
| hovermode="x unified", | |
| xaxis_title="Signed severity (− sharpen / + degrade)", | |
| yaxis_title={ | |
| "dice": "Dice", | |
| "hd95": "HD95 (vox)", | |
| "hd100": "HD100 (vox)", | |
| "asd": "ASD (vox)", | |
| }.get(_cmp_metric, _cmp_metric), | |
| legend=dict(orientation="h", yanchor="bottom", y=1.02, x=0), | |
| margin=dict(l=20, r=20, t=20, b=20), | |
| ) | |
| st.plotly_chart(_fig_cmp, use_container_width=True) | |
| else: | |
| st.info("Install plotly for interactive charts.") | |
| with tab_dl_cmp: | |
| st.markdown("#### Model A Downloads") | |
| col1, col2 = st.columns(2) | |
| with col1: | |
| st.download_button( | |
| "Download A — JSON", | |
| data=_ev_a["card_json"], | |
| file_name=f"card_A_{cmp_a_dataset}_{cmp_a_arch}_seed{cmp_a_seed}.json", | |
| mime="application/json", | |
| use_container_width=True, | |
| ) | |
| with col2: | |
| st.download_button( | |
| "Download A — CSV", | |
| data=_ev_a["df"].to_csv(index=False).encode(), | |
| file_name=f"metrics_A_{cmp_a_dataset}_{cmp_a_arch}_seed{cmp_a_seed}.csv", | |
| mime="text/csv", | |
| use_container_width=True, | |
| ) | |
| st.markdown("#### Model B Downloads") | |
| col3, col4 = st.columns(2) | |
| with col3: | |
| st.download_button( | |
| "Download B — JSON", | |
| data=_ev_b["card_json"], | |
| file_name=f"card_B_{cmp_b_dataset}_{cmp_b_arch}_seed{cmp_b_seed}.json", | |
| mime="application/json", | |
| use_container_width=True, | |
| ) | |
| with col4: | |
| st.download_button( | |
| "Download B — CSV", | |
| data=_ev_b["df"].to_csv(index=False).encode(), | |
| file_name=f"metrics_B_{cmp_b_dataset}_{cmp_b_arch}_seed{cmp_b_seed}.csv", | |
| mime="text/csv", | |
| use_container_width=True, | |
| ) | |
| st.markdown("#### Print / Save as PDF") | |
| _print_as_pdf_button("Print comparison page as PDF") | |
| st.stop() | |
| # =========================================================================== | |
| # MODE 3: UPLOAD YOUR MODEL | |
| # =========================================================================== | |
| st.markdown( | |
| """<div class="segroc-hero"> | |
| <h2>Upload Your Model — Single Sample Evaluation</h2> | |
| <p>Run a live robustness sweep on one image using your own MONAI checkpoint.</p> | |
| </div>""", | |
| unsafe_allow_html=True, | |
| ) | |
| st.markdown( | |
| """<div class="segroc-disclaimer"> | |
| <b>⚠️ MONAI models only.</b> This mode supports MONAI-framework | |
| architectures: UNet, Attention UNet, UNETR, SwinUNETR, and VNet. | |
| Upload a <code>.pt</code> / <code>.ckpt</code> file; the architecture | |
| is inferred from the weight tensor shapes. You are responsible for | |
| confirming the correct architecture in the sidebar before running. | |
| Ground-truth masks are optional — if omitted, metrics are computed | |
| relative to the unperturbed baseline prediction. | |
| </div>""", | |
| unsafe_allow_html=True, | |
| ) | |
| if not run_btn and "upload_eval_results" not in st.session_state: | |
| st.info("Configure the sidebar and press **Run Evaluation** to start.") | |
| st.stop() | |
| # ── Run evaluation on button press ───────────────────────────────────────── | |
| if run_btn: | |
| with st.spinner("Preparing model…"): | |
| if ckpt_bytes is not None: | |
| try: | |
| model = load_model_from_bytes( | |
| ckpt_bytes, final_arch, final_in_ch, final_out_ch, final_spatial | |
| ) | |
| actual_out = model.out_channels | |
| except RuntimeError as e: | |
| st.error( | |
| f"Failed to load checkpoint with architecture '{final_arch}'. " | |
| f"Verify the architecture matches the checkpoint. " | |
| f"Error: {str(e)[:200]}…" | |
| ) | |
| st.stop() | |
| else: | |
| model = get_demo_model(final_arch, final_in_ch, final_out_ch, final_spatial) | |
| actual_out = final_out_ch | |
| gt_mask: torch.Tensor | None = None | |
| with st.spinner("Loading image…"): | |
| if npz_file is not None: | |
| _npz_bytes = npz_file.getvalue() | |
| image, _display_np, _npz_lbl = load_npz_sample(_npz_bytes, npz_sample_idx) | |
| vol_affine = np.eye(4) | |
| if _npz_lbl is not None: | |
| gt_mask = torch.from_numpy(_npz_lbl).long() | |
| elif png_file is not None: | |
| image, _display_np = load_png_bytes( | |
| png_file.getvalue(), filename=png_file.name | |
| ) | |
| vol_affine = np.eye(4) | |
| if png_label_file is not None: | |
| label_arr = load_png_label_bytes( | |
| png_label_file.getvalue(), filename=png_label_file.name | |
| ) | |
| gt_mask = torch.from_numpy(label_arr).long() | |
| elif nii_file is not None: | |
| image, _, vol_affine = load_nifti_bytes( | |
| nii_file.getvalue(), filename=nii_file.name | |
| ) | |
| else: | |
| image, _syn_gt = make_synthetic_volume() | |
| gt_mask = _syn_gt | |
| vol_affine = np.eye(4) | |
| st.caption("Using synthetic 64³ ellipsoid (demo).") | |
| if gt_file is not None: | |
| with st.spinner("Loading ground-truth label…"): | |
| gt_mask = load_label_bytes(gt_file.getvalue(), filename=gt_file.name) | |
| img_spatial = image.dim() - 2 | |
| if img_spatial != final_spatial: | |
| st.error( | |
| f"Dimension mismatch: image has **{img_spatial}D** but model is " | |
| f"configured for **{final_spatial}D**." | |
| ) | |
| st.stop() | |
| img_ch = image.shape[1] | |
| if img_ch != final_in_ch: | |
| st.error( | |
| f"Channel mismatch: image has **{img_ch}** channel(s) but model " | |
| f"expects **{final_in_ch}**." | |
| ) | |
| st.stop() | |
| if gt_mask is not None: | |
| expected_sp = tuple(image.shape[2:]) | |
| if tuple(gt_mask.shape) != expected_sp: | |
| st.error( | |
| f"Label shape {tuple(gt_mask.shape)} doesn't match image " | |
| f"spatial shape {expected_sp}." | |
| ) | |
| st.stop() | |
| if gt_mask is None: | |
| st.info( | |
| "No ground-truth label — segmentation metrics will be computed " | |
| "relative to the **unperturbed baseline prediction**." | |
| ) | |
| is_2d = image.dim() == 4 | |
| severity_schedule = build_severity_schedule( | |
| n_sharp=int(n_sharp), | |
| n_degrade=int(n_degrade), | |
| max_sharp=float(max_sharp), | |
| max_degrade=float(max_degrade), | |
| ) | |
| total_levels = len(severity_schedule) | |
| baseline_idx = severity_schedule.index(0.0) | |
| perturbed_vols: list[tuple[float, np.ndarray]] = [] | |
| pred_vols: list[np.ndarray] = [] | |
| dim_label = "2-D" if is_2d else "3-D" | |
| _prog = st.progress(0, text=f"Running {dim_label} inference…") | |
| for i, sev in enumerate(severity_schedule): | |
| perturbed = get_perturbed_image( | |
| image, sev, artifact=artifact, modality=modality, seed=int(seed) | |
| ) | |
| pred = run_inference(model, perturbed) | |
| perturbed_sq = perturbed.squeeze(0) | |
| if is_2d: | |
| disp = perturbed_sq.mean(0).cpu().numpy()[np.newaxis] | |
| pred_disp = pred.cpu().numpy()[np.newaxis] | |
| else: | |
| disp = perturbed_sq.squeeze(0).cpu().numpy() | |
| pred_disp = pred.cpu().numpy() | |
| perturbed_vols.append((sev, disp)) | |
| pred_vols.append(pred_disp) | |
| _prog.progress( | |
| (i + 1) / total_levels, | |
| text=f"Inference {i+1}/{total_levels} (sev={sev:+.2f})", | |
| ) | |
| _prog.empty() | |
| if gt_mask is not None: | |
| reference_tensor = gt_mask | |
| reference_label = "vs. GT" | |
| else: | |
| reference_tensor = torch.from_numpy(pred_vols[baseline_idx]).long() | |
| if is_2d: | |
| reference_tensor = reference_tensor.squeeze(0) | |
| reference_label = "vs. baseline pred" | |
| rows: list[dict] = [] | |
| _prog2 = st.progress(0, text="Computing metrics…") | |
| for i, sev in enumerate(severity_schedule): | |
| pred_for_metric = torch.from_numpy(pred_vols[i]).long() | |
| if is_2d: | |
| pred_for_metric = pred_for_metric.squeeze(0) | |
| metrics = compute_all_metrics(pred_for_metric, reference_tensor, actual_out) | |
| direction = "sharp" if sev < 0 else ("base" if sev == 0.0 else "degrade") | |
| row: dict = {"severity": round(sev, 4), "direction": direction} | |
| row.update({k: round(v, 5) for k, v in metrics.items()}) | |
| rows.append(row) | |
| _prog2.progress((i + 1) / total_levels, text=f"Metrics {i+1}/{total_levels}") | |
| _prog2.empty() | |
| baseline_vol = perturbed_vols[baseline_idx][1] | |
| for i, (sev, pt_np) in enumerate(perturbed_vols): | |
| rows[i]["rmse"] = round(float(np.sqrt(np.mean((pt_np - baseline_vol) ** 2))), 6) | |
| rows[i]["psnr"] = round(psnr(pt_np, baseline_vol), 4) | |
| rows[i]["ssim"] = round(ssim_score(pt_np, baseline_vol), 6) | |
| baseline_row = rows[baseline_idx] | |
| degrade_rows = [r for r in rows if r["direction"] == "degrade"] | |
| def _rood(key: str, higher: bool) -> dict[str, float]: | |
| base_val = baseline_row.get(key, float("nan")) | |
| lvl_vals = [r.get(key, float("nan")) for r in degrade_rows] | |
| return { | |
| f"wm_{key}": wm_metric_t(base_val, lvl_vals), | |
| f"m_ddeg_{key}": m_ddeg_t(base_val, lvl_vals, higher_is_better=higher), | |
| } | |
| rood_metrics: dict[str, float] = {} | |
| rood_metrics.update(_rood("dice", higher=True)) | |
| if _SCIPY and "hd95" in baseline_row: | |
| rood_metrics.update(_rood("hd95", higher=False)) | |
| else: | |
| rood_metrics["wm_hd95"] = float("nan") | |
| rood_metrics["m_ddeg_hd95"] = float("nan") | |
| rood_out = { | |
| "wm_dsc": rood_metrics.get("wm_dice", float("nan")), | |
| "wm_hd95": rood_metrics.get("wm_hd95", float("nan")), | |
| "m_ddeg": rood_metrics.get("m_ddeg_dice", float("nan")), | |
| "m_ddeg_hd": rood_metrics.get("m_ddeg_hd95", float("nan")), | |
| } | |
| df_up = pd.DataFrame(rows) | |
| gt_np_viewer: np.ndarray | None = None | |
| if gt_mask is not None: | |
| gt_np_viewer = gt_mask.cpu().numpy() | |
| if gt_np_viewer.ndim == 2: | |
| gt_np_viewer = gt_np_viewer[np.newaxis] | |
| st.session_state["upload_eval_results"] = { | |
| "rows": rows, | |
| "df": df_up, | |
| "perturbed_vols": perturbed_vols, | |
| "pred_vols": pred_vols, | |
| "baseline_idx": baseline_idx, | |
| "total_levels": total_levels, | |
| "severity_schedule": severity_schedule, | |
| "actual_out": actual_out, | |
| "gt_mask_np": gt_np_viewer, | |
| "reference_label": reference_label, | |
| "artifact": artifact, | |
| "vol_affine": vol_affine, | |
| "rood": rood_out, | |
| "is_2d": is_2d, | |
| } | |
| st.success( | |
| f"Evaluation complete — {total_levels} levels. Metrics computed {reference_label}." | |
| ) | |
| # ── Render upload eval results ────────────────────────────────────────────── | |
| _up_ev = st.session_state["upload_eval_results"] | |
| rows = _up_ev["rows"] | |
| df = _up_ev["df"] | |
| perturbed_vols = _up_ev["perturbed_vols"] | |
| pred_vols = _up_ev["pred_vols"] | |
| baseline_idx = _up_ev["baseline_idx"] | |
| total_levels = _up_ev["total_levels"] | |
| severity_schedule = _up_ev["severity_schedule"] | |
| actual_out = _up_ev["actual_out"] | |
| gt_np = _up_ev["gt_mask_np"] | |
| reference_label = _up_ev["reference_label"] | |
| artifact = _up_ev["artifact"] | |
| vol_affine = _up_ev["vol_affine"] | |
| rood = _up_ev["rood"] | |
| is_2d = _up_ev.get("is_2d", False) | |
| available_metrics: list[str] = ( | |
| ["dice"] + (["hd95", "hd100", "asd"] if _SCIPY else []) + ["rmse", "psnr", "ssim"] | |
| ) | |
| available_metrics = [m for m in available_metrics if m in df.columns] | |
| _D, _H, _W = perturbed_vols[0][1].shape | |
| _vol_label = "2-D" if is_2d else "3-D" | |
| # ── Summary KPIs ───────────────────────────────────────────────────────────── | |
| baseline_row_data = ( | |
| df[df["severity"] == 0.0].iloc[0] if 0.0 in df["severity"].values else None | |
| ) | |
| c1, c2, c3, c4 = st.columns(4) | |
| with c1: | |
| _dice_val = ( | |
| f"{baseline_row_data['dice']:.4f}" | |
| if baseline_row_data is not None | |
| and not np.isnan(baseline_row_data.get("dice", float("nan"))) | |
| else "N/A" | |
| ) | |
| _dice_lbl = f"Baseline Dice ({reference_label})" | |
| st.metric(_dice_lbl, _dice_val) | |
| with c2: | |
| _wm = rood.get("wm_dsc", float("nan")) | |
| st.metric("wmDSCt (↑ better)", "N/A" if np.isnan(_wm) else f"{_wm:.4f}") | |
| with c3: | |
| _md = rood.get("m_ddeg", float("nan")) | |
| st.metric("mDDegt (↓ better)", "N/A" if np.isnan(_md) else f"{_md:.4f}") | |
| with c4: | |
| _wh = rood.get("wm_hd95", float("nan")) | |
| st.metric("wmHD95t (↓ better)", "N/A" if np.isnan(_wh) else f"{_wh:.4f}") | |
| # ── Tabs ────────────────────────────────────────────────────────────────────── | |
| tab_curve, tab_images, tab_seg, tab_table = st.tabs( | |
| [ | |
| "Metrics", | |
| f"Image Quality ({_vol_label})", | |
| f"Segmentation ({_vol_label})", | |
| "Metrics Table", | |
| ] | |
| ) | |
| # ── Tab: Metrics ───────────────────────────────────────────────────────────── | |
| with tab_curve: | |
| if not _PLOTLY: | |
| st.line_chart(df.set_index("severity")[available_metrics]) | |
| else: | |
| active_metrics = st.multiselect( | |
| "Metrics to display", | |
| available_metrics, | |
| default=[m for m in ["dice", "ssim"] if m in available_metrics], | |
| format_func=lambda k: _METRIC_META.get(k, {}).get("label", k), | |
| key="up_metrics_sel", | |
| ) | |
| if not active_metrics: | |
| st.warning("Select at least one metric.") | |
| else: | |
| fig_curve = make_subplots(specs=[[{"secondary_y": True}]]) | |
| for mk in active_metrics: | |
| meta = _METRIC_META.get(mk, {}) | |
| sec = meta.get("axis", "left") == "right" | |
| fig_curve.add_trace( | |
| go.Scatter( | |
| x=df["severity"].tolist(), | |
| y=df[mk].tolist(), | |
| mode="lines+markers", | |
| name=meta.get("label", mk), | |
| line=dict(color=meta.get("colour", "#334155"), width=2.3), | |
| marker=dict(size=6), | |
| ), | |
| secondary_y=sec, | |
| ) | |
| fig_curve.add_vline( | |
| x=0.0, line_dash="dash", line_color="rgba(100,116,139,0.6)" | |
| ) | |
| fig_curve.update_layout( | |
| height=480, | |
| template="plotly_white", | |
| hovermode="x unified", | |
| legend=dict(orientation="h", yanchor="bottom", y=1.02, x=0), | |
| margin=dict(l=20, r=20, t=20, b=20), | |
| ) | |
| fig_curve.update_xaxes(title_text="Signed severity (− sharpen / + degrade)") | |
| fig_curve.update_yaxes(title_text="Segmentation / SSIM", secondary_y=False) | |
| fig_curve.update_yaxes(title_text="HD / RMSE / PSNR", secondary_y=True) | |
| st.plotly_chart(fig_curve, use_container_width=True) | |
| # ── Tab: Image Quality ──────────────────────────────────────────────────────── | |
| with tab_images: | |
| show_seg_overlay = st.toggle( | |
| "Show segmentation overlay", value=True, key="up_seg_overlay_imgs" | |
| ) | |
| if _D == 1: | |
| overview_sev = st.select_slider( | |
| "Severity level", | |
| options=[round(s, 4) for s in severity_schedule], | |
| value=0.0, | |
| key="up_sev_slider_imgs", | |
| ) | |
| overview_idx = [round(s, 4) for s in severity_schedule].index( | |
| round(overview_sev, 4) | |
| ) | |
| _, vol_disp = perturbed_vols[overview_idx] | |
| pred_disp = pred_vols[overview_idx] | |
| _render_three_planes( | |
| vol_disp, | |
| pred_disp if show_seg_overlay else None, | |
| gt_np, | |
| height=280, | |
| key_prefix=f"up_img_{overview_idx}", | |
| ) | |
| else: | |
| overview_sev = st.select_slider( | |
| "Severity level", | |
| options=[round(s, 4) for s in severity_schedule], | |
| value=0.0, | |
| key="up_sev_slider_3d", | |
| ) | |
| overview_idx = [round(s, 4) for s in severity_schedule].index( | |
| round(overview_sev, 4) | |
| ) | |
| ax_i = st.slider("Axial slice", 0, max(0, _D - 1), _D // 2, key="up_ax") | |
| cor_i = st.slider("Coronal slice", 0, max(0, _H - 1), _H // 2, key="up_cor") | |
| sag_i = st.slider("Sagittal slice", 0, max(0, _W - 1), _W // 2, key="up_sag") | |
| _, vol_disp = perturbed_vols[overview_idx] | |
| pred_disp = pred_vols[overview_idx] | |
| _render_three_planes( | |
| vol_disp, | |
| pred_disp if show_seg_overlay else None, | |
| gt_np, | |
| height=220, | |
| key_prefix=f"up_3d_{overview_idx}", | |
| ax_idx=ax_i, | |
| cor_idx=cor_i, | |
| sag_idx=sag_i, | |
| ) | |
| # ── Tab: Segmentation ───────────────────────────────────────────────────────── | |
| with tab_seg: | |
| seg_col_a, seg_col_b = st.columns(2) | |
| with seg_col_a: | |
| sev_a = st.select_slider( | |
| "Severity A", | |
| options=[round(s, 4) for s in severity_schedule], | |
| value=0.0, | |
| key="up_sev_a", | |
| ) | |
| with seg_col_b: | |
| sev_b = st.select_slider( | |
| "Severity B", | |
| options=[round(s, 4) for s in severity_schedule], | |
| value=round(severity_schedule[-1], 4), | |
| key="up_sev_b", | |
| ) | |
| sev_a_idx = [round(s, 4) for s in severity_schedule].index(round(sev_a, 4)) | |
| sev_b_idx = [round(s, 4) for s in severity_schedule].index(round(sev_b, 4)) | |
| col_left, col_right = st.columns(2) | |
| with col_left: | |
| st.caption(f"Severity {sev_a:+.2f}") | |
| _, vol_a = perturbed_vols[sev_a_idx] | |
| _render_three_planes( | |
| vol_a, | |
| pred_vols[sev_a_idx], | |
| gt_np, | |
| height=220, | |
| key_prefix=f"up_seg_a_{sev_a_idx}", | |
| ) | |
| with col_right: | |
| st.caption(f"Severity {sev_b:+.2f}") | |
| _, vol_b = perturbed_vols[sev_b_idx] | |
| _render_three_planes( | |
| vol_b, | |
| pred_vols[sev_b_idx], | |
| gt_np, | |
| height=220, | |
| key_prefix=f"up_seg_b_{sev_b_idx}", | |
| ) | |
| # ── Tab: Metrics Table ──────────────────────────────────────────────────────── | |
| with tab_table: | |
| st.dataframe(df, use_container_width=True, hide_index=True) | |
| st.download_button( | |
| "Download metrics CSV", | |
| data=df.to_csv(index=False).encode(), | |
| file_name="upload_eval_metrics.csv", | |
| mime="text/csv", | |
| ) | |
| if _NIB and vol_affine is not None: | |
| baseline_pred_np = pred_vols[baseline_idx] | |
| if is_2d: | |
| baseline_pred_np = baseline_pred_np.squeeze(0) | |
| nii_bytes_out = save_nifti_bytes(baseline_pred_np.astype(np.int16), vol_affine) | |
| if nii_bytes_out: | |
| st.download_button( | |
| "Download baseline prediction (NIfTI)", | |
| data=nii_bytes_out, | |
| file_name="baseline_prediction.nii.gz", | |
| mime="application/gzip", | |
| ) | |