Source code for spacr.classifier_evaluation

"""Classifier evaluation, calibration, and split-leakage diagnostics.

The module is model-agnostic: it consumes labels, probabilities, fold ids,
and sample paths. Deep-learning CV and future classical-ML pipelines can
therefore write the same evaluation bundle and the Qt workbench can display
one stable artifact format.
"""
from __future__ import annotations

from dataclasses import asdict, dataclass, field
import hashlib
import json
import math
import os
from pathlib import Path
import re
from typing import Any, Dict, Iterable, List, Mapping, Optional, Sequence, Tuple

import numpy as np
import pandas as pd
from sklearn.metrics import (
    accuracy_score,
    confusion_matrix,
    f1_score,
    log_loss,
    precision_score,
    recall_score,
)

from .figures.style import figure_style, theme_target


[docs] EVALUATION_FILES = { "summary": "summary.json", "predictions": "oof_predictions.csv", "confusion_counts": "confusion_counts.csv", "confusion_normalized": "confusion_normalized.csv", "per_plate": "per_plate_metrics.csv", "calibration": "calibration.csv", "leakage": "leakage.json", "manifest": "evaluation_manifest.json", "confusion_figure": "confusion_matrix.png", "calibration_figure": "calibration.png", }
"""Stable file names produced by :func:`write_evaluation_bundle`.""" SPLIT_LEVELS: Tuple[str, ...] = ("cell", "field", "well", "plate") _SPLIT_COLUMNS: Dict[str, Tuple[str, ...]] = { "cell": (), "field": ("plateID", "rowID", "columnID", "fieldID"), "well": ("plateID", "rowID", "columnID"), "plate": ("plateID",), } @dataclass(frozen=True)
[docs] class SplitReport: """Provenance and realised sizes for one train/test split. :param group_by: canonical isolation unit used for the split, from objects through fields, wells, and plates. :param requested_fraction: held-out object fraction requested by the caller. :param cell_fraction: realised share of objects assigned to the test side. :param group_fraction: realised share of distinct groups assigned to test. :param train_cells: number of object rows used for fitting. :param test_cells: number of object rows held out for evaluation. :param train_groups: number of distinct groups used for fitting. :param test_groups: number of distinct groups held out for evaluation. :param total_groups: distinct groups across both sides of the split. :param rule: human-readable algorithm and isolation guarantee that produced the realised split. """ group_by: str requested_fraction: float cell_fraction: float group_fraction: float train_cells: int test_cells: int train_groups: int test_groups: int total_groups: int rule: str
[docs] def to_dict(self) -> Dict[str, Any]: """Return JSON-safe split provenance for model cards/settings.""" return asdict(self)
[docs] def summary(self) -> str: """Describe both whole-group and object-level holdout costs.""" unit = self.group_by if self.group_by != "cell" else "object" return ( f"split_by={self.group_by}: held out {self.test_groups} of " f"{self.total_groups} {unit} group(s) ({self.group_fraction:.1%}) " f"and {self.test_cells} of {self.train_cells + self.test_cells} " f"cells ({self.cell_fraction:.1%}); requested " f"{self.requested_fraction:.1%}. {self.rule}" )
[docs] def normalize_split_level(group_by: Any) -> str: """Return a canonical split level; legacy ``none`` means ``cell``. :param group_by: requested split level: ``cell``, ``field``, ``well`` or ``plate`` (exact, lower-case). ``None``, ``False``, ``'none'`` and ``'off'`` map to ``cell``; anything else raises :class:`ValueError`. """ if group_by in (None, False, "none", "off"): return "cell" level = str(group_by) if level not in SPLIT_LEVELS: choices = ", ".join(SPLIT_LEVELS) raise ValueError( f"Unknown group_by/cv_group_by split level {group_by!r}; use one of " f"{choices}. Names " "are exact and lower-case. 'none' remains an alias for 'cell'." ) return level
[docs] def split_columns_for(group_by: Any, columns: Sequence[str], table: str = "data") -> Tuple[str, List[str]]: """Resolve a split level to complete metadata columns or refuse it. ``prcfo`` and crop filenames are handled by :func:`split_group_values`; this helper describes the direct-column route used by measurement tables. Partial keys are never accepted because, for example, ``columnID='c1'`` is not a well identity across rows and plates. :param group_by: split level, normalised by :func:`normalize_split_level`. :param columns: column names available in the table; every identity column the level needs must be present or :class:`ValueError` is raised. """ level = normalize_split_level(group_by) wanted = list(_SPLIT_COLUMNS[level]) if not wanted: return level, [] missing = [column for column in wanted if column not in columns] if missing: raise ValueError( f"Cannot split {table} by {level}: the complete identity needs " f"{wanted}, but {missing} {'is' if len(missing) == 1 else 'are'} " "missing. Use a frame with acquisition metadata, a valid prcfo, " "or explicitly choose a finer split level such as 'cell'." ) return level, wanted
def _groups_from_prcfo(values: Sequence[Any], level: str, table: str) -> np.ndarray: """Parse strict measurement identities from canonical object keys.""" from . import schema groups: List[str] = [] for position, value in enumerate(values): try: identity = schema.parse_prcfo(value) except (TypeError, ValueError) as exc: raise ValueError( f"Cannot split {table} by {level}: prcfo at row {position} " f"is not a valid object identity ({value!r}): {exc}" ) from exc parts = identity.to_dict() wanted = _SPLIT_COLUMNS[level] groups.append("\x1f".join(str(parts[column]) for column in wanted)) return np.asarray(groups, dtype=object)
[docs] def split_group_values(*, group_by: Any = "well", frame: Optional[pd.DataFrame] = None, paths: Optional[Sequence[Any]] = None, table: str = "data") -> Tuple[str, np.ndarray]: """Build one strict group id per row from metadata, ``prcfo``, or paths. For grouped levels every identity must be verifiable. Inventing singleton ids for unparseable rows would make a leaking random split look grouped. """ level = normalize_split_level(group_by) if frame is None and paths is None: raise ValueError("split_group_values needs a frame or paths") n = len(frame) if frame is not None else len(paths or ()) if level == "cell": if frame is not None and "prcfo" in frame.columns: values = frame["prcfo"].astype(str) if frame["prcfo"].isna().any() or values.str.strip().eq("").any(): raise ValueError( f"Cannot split {table} by cell: at least one row has no " "object identity.") return level, values.to_numpy(dtype=object) if frame is not None and frame.index.name == "prcfo": values = pd.Series(frame.index.astype(str)) if values.str.strip().eq("").any(): raise ValueError( f"Cannot split {table} by cell: at least one row has no " "object identity.") return level, values.to_numpy(dtype=object) if paths is not None: return level, np.asarray( [augmentation_family(path) for path in paths], dtype=object) return level, np.arange(n, dtype=np.int64).astype(object) wanted = list(_SPLIT_COLUMNS[level]) if frame is not None and all(column in frame.columns for column in wanted): metadata = frame[wanted] missing = metadata.isna() | metadata.astype(str).apply( lambda column: column.str.strip().eq("")) if bool(missing.to_numpy().any()): row, column = np.argwhere(missing.to_numpy())[0] raise ValueError( f"Cannot split {table} by {level}: row {int(row)} has no " f"{wanted[int(column)]} identity. A random fallback would " "report an optimistic score as grouped." ) return level, metadata.astype(str).agg("\x1f".join, axis=1).to_numpy() if frame is not None: if "prcfo" in frame.columns: return level, _groups_from_prcfo(frame["prcfo"].tolist(), level, table) if frame.index.name == "prcfo": return level, _groups_from_prcfo(frame.index.tolist(), level, table) if paths is None and "png_path" in frame.columns: paths = frame["png_path"].tolist() if paths is not None: groups = [] for position, path in enumerate(paths): identity = sample_identity(path) encoded = augmentation_family(path).split("_") value = identity.get(level, "") if len(encoded) < 4 or not value: raise ValueError( f"Cannot split {table} by {level}: path at row " f"{position} ({path!r}) does not encode a {level}. " "spaCR crop names must encode plate_well_field before " "the object; choose 'cell' explicitly only if leakage " "between sibling crops is intended." ) groups.append(value) return level, np.asarray(groups, dtype=object) missing = [column for column in wanted if frame is None or column not in frame.columns] raise ValueError( f"Cannot split {table} by {level}: missing identity columns {missing}, " "and no valid prcfo or crop path is available. A grouped design is " "therefore unverifiable and will not be replaced by a random split." )
class _GroupedSplitImpossible(ValueError): """Raised when the labelled groups cannot give a grouped holdout. Every labelled object comes from one group, or no arrangement of the groups puts every class on both sides. A ``ValueError``, so every caller that already catches one is unchanged; the subclass lets a caller whose product is not the held-out score tell this refusal apart from a missing label or a bad setting. """
[docs] def grouped_split(groups: Sequence[Any], labels: Sequence[Any], holdout: float, seed: int = 0, *, group_by: Any = "well", hold_out_groups: Optional[Sequence[Any]] = None ) -> Tuple[np.ndarray, np.ndarray, SplitReport]: """Return a stratified holdout while keeping named groups intact. A grouped design is refused when either side cannot contain every class. This is intentionally stricter than silently scoring a model on siblings of its training rows or on a holdout that contains only one class. :param groups: group identifier for every labelled object. All members of one group remain on the same side of the split. :param labels: class label for every object, aligned to ``groups``. :param holdout: requested test fraction, strictly between zero and one. :param hold_out_groups: groups that go to the TEST side whatever the fraction says. This is what `holdout_plate` is: cross-validation splits within the data it is given, so a model can learn the plate rather than the phenotype and every number it reports will look fine. Naming a plate here trains without it and scores on it, which is the one number that says whether a classifier generalises. The class check still applies: a named holdout that leaves either side without every class is refused, for the same reason a random one is. """ from sklearn.model_selection import (GroupShuffleSplit, StratifiedGroupKFold, train_test_split) level = normalize_split_level(group_by) y = np.asarray(labels) group_values = np.asarray(groups, dtype=object) fraction = float(holdout) if not np.isfinite(fraction) or not 0.0 < fraction < 1.0: raise ValueError("holdout must be a finite fraction strictly between 0 and 1") if len(y) == 0: raise ValueError( "there are no labelled objects to split, so a classifier cannot " "be trained or scored. This usually means the control values " "matched no rows, or that filtering removed every row before the " "split. Check that positive_control_id and negative_control_id " "name " "values present in the control column, and that any measurement " "filters still leave objects behind.") if len(group_values) != len(y): raise ValueError( "group-aware splitting requires one group per label; the split " f"has {len(group_values)} groups for {len(y)} labels") if len(y) < 2: raise ValueError("a train/test split needs at least two labelled cells") classes = np.unique(y) if hold_out_groups: wanted = {str(g).strip() for g in hold_out_groups if str(g).strip()} as_text = np.array([str(g) for g in group_values]) test_mask = np.isin(as_text, list(wanted)) if not test_mask.any(): raise ValueError( f"none of the held-out {level}(s) {sorted(wanted)} appear in " f"the data; it has {sorted(set(as_text))[:8]}") if test_mask.all(): raise ValueError( f"the held-out {level}(s) {sorted(wanted)} are ALL of the " f"data, so there is nothing left to train on") train_idx = np.where(~test_mask)[0] test_idx = np.where(test_mask)[0] for side, name in ((train_idx, "training"), (test_idx, "held-out")): missing = set(classes) - set(np.unique(y[side])) if missing: raise ValueError( f"holding out {level}(s) {sorted(wanted)} leaves the " f"{name} side without class(es) {sorted(missing)}, so the " f"score would not mean what it says") total_groups = len(set(as_text)) test_groups = len(set(as_text[test_mask])) report = SplitReport( group_by=level, requested_fraction=float(fraction), cell_fraction=float(test_mask.sum()) / float(len(y)), group_fraction=float(test_groups) / float(max(1, total_groups)), train_cells=int(len(train_idx)), test_cells=int(len(test_idx)), train_groups=int(total_groups - test_groups), test_groups=int(test_groups), total_groups=int(total_groups), rule=f"held out by name: {', '.join(sorted(wanted))}") return train_idx, test_idx, report indices = np.arange(len(y)) distinct = np.unique(group_values.astype(str)) protects_groups = level != "cell" or len(distinct) < len(y) if not protects_groups: counts = pd.Series(y).value_counts() if len(counts) > 1 and int(counts.min()) < 2: raise ValueError( "A cell split cannot put every class in both train and test: " f"class counts are {counts.to_dict()}. Add another labelled " "cell in the rare class." ) stratify = y if int(counts.min()) >= 2 else None train_idx, test_idx = train_test_split( indices, test_size=fraction, random_state=int(seed), stratify=stratify, ) rule = "stratified random split of objects; sibling cells may cross" else: if any(value is None or str(value).strip() == "" for value in group_values): raise ValueError( f"Cannot split by {level}: at least one cell has no group " "identity, so independence cannot be verified." ) if len(distinct) < 2: where = distinct[0] if len(distinct) else "unknown" unit = "object" if level == "cell" else level raise _GroupedSplitImpossible( f"A {unit}-grouped held-out split is impossible: every " f"labelled cell comes from one {unit} ({where}). A random " "cell split would only measure how well the model memorised " f"this {unit}, not whether it transfers." ) candidates: List[Tuple[np.ndarray, np.ndarray, str]] = [] requested_splits = max( 2, min(int(round(1.0 / fraction)), len(distinct), 20)) split_counts = sorted({ requested_splits, *range(2, min(5, len(distinct)) + 1), }) for n_splits in split_counts: try: splitter = StratifiedGroupKFold( n_splits=n_splits, shuffle=True, random_state=int(seed)) for train, test in splitter.split(indices, y, group_values): candidates.append((train, test, f"StratifiedGroupKFold({n_splits}) over " f"{len(distinct)} {level} groups")) except ValueError: continue splitter = GroupShuffleSplit( n_splits=min(256, max(32, len(distinct) * 4)), test_size=fraction, random_state=int(seed)) try: for train, test in splitter.split(indices, y, group_values): candidates.append((train, test, f"GroupShuffleSplit over {len(distinct)} {level} groups")) except ValueError: pass complete = [candidate for candidate in candidates if set(np.unique(y[candidate[0]])) == set(classes) and set(np.unique(y[candidate[1]])) == set(classes)] if not complete: per_class = { str(label): int(len(np.unique(group_values[y == label]))) for label in classes } raise _GroupedSplitImpossible( f"A leakage-safe {level}-grouped split cannot put every " "class in both train and test. Independent groups per class: " f"{per_class}. Add independent {level}s, choose a finer " "level, or collect another class-bearing group; a random " "fallback would report memorisation as transfer." ) target = fraction * len(y) train_idx, test_idx, rule = min( complete, key=lambda candidate: ( abs(len(np.unique(group_values[candidate[1]])) / len(distinct) - fraction), abs(len(candidate[1]) - target), ), ) train_groups = set(group_values[train_idx].astype(str)) test_groups = set(group_values[test_idx].astype(str)) if train_groups & test_groups: raise RuntimeError(f"{level} groups crossed the train/test boundary") train_idx = np.sort(np.asarray(train_idx, dtype=int)) test_idx = np.sort(np.asarray(test_idx, dtype=int)) total_groups = len(np.unique(group_values.astype(str))) test_groups_n = len(np.unique(group_values[test_idx].astype(str))) if level == "cell" and protects_groups: rule += ( "; repeated rows of one object stay together; sibling cells may " "cross") report = SplitReport( group_by=level, requested_fraction=fraction, cell_fraction=len(test_idx) / len(y), group_fraction=test_groups_n / total_groups, train_cells=len(train_idx), test_cells=len(test_idx), train_groups=total_groups - test_groups_n, test_groups=test_groups_n, total_groups=total_groups, rule=rule + ("; no group appears on both sides" if protects_groups else ""), ) return train_idx, test_idx, report
[docs] class LeakageError(ValueError): """Raised when related samples cross a protected split boundary."""
@dataclass
[docs] class LeakageReport: """Overlap counts and examples for one train/validation boundary. :param group_by: protected split level (``cell``, ``field``, ``well``, or ``plate``); the legacy ``none`` spelling is normalized to ``cell``. :param train_samples: number of training paths. :param validation_samples: number of validation paths. :param overlap_counts: overlap count at each identity level. :param examples: up to ten shared identities per level. :param split_name: caller-supplied label for the audited boundary. :param critical_levels: levels that invalidate the requested split. :param warnings: non-fatal caveats. :param unverifiable_counts: samples lacking a requested identity or content hash, counted by the level that could not be verified. :param hash_errors: up to twenty file-specific failures from optional byte-content hashing. """ group_by: str train_samples: int validation_samples: int overlap_counts: Dict[str, int] examples: Dict[str, List[str]] split_name: str = "" critical_levels: List[str] = field(default_factory=list) warnings: List[str] = field(default_factory=list) unverifiable_counts: Dict[str, int] = field(default_factory=dict) hash_errors: List[str] = field(default_factory=list) @property
[docs] def passed(self) -> bool: """Return True when no protected identity crosses the boundary.""" return not self.critical_levels
[docs] def to_dict(self) -> Dict[str, Any]: """Return a JSON-serializable report.""" data = asdict(self) data["passed"] = self.passed return data
def _stem(path: Any) -> str: """Return a normalized basename without its final extension.""" return os.path.splitext(os.path.basename(str(path)))[0] _AUGMENT_SUFFIX = re.compile( r"(?i)(?:[_-](?:aug(?:ment)?\d*|rot(?:ate)?(?:90|180|270)|" r"flip(?:ped)?[_-]?[hv]|horizontal|vertical))+$" )
[docs] def augmentation_family(path: Any) -> str: """Return the original-sample family after removing augmentation suffixes. In-memory spaCR augmentations retain the exact source filename, while older exported datasets use suffixes such as ``_aug3``, ``_rot90`` or ``_flip_h``. Both forms collapse to one family. :param path: crop path or file name; only its basename without the final extension is used, with trailing augmentation suffixes stripped repeatedly. """ stem = _stem(path) previous = None while previous != stem: previous = stem stem = _AUGMENT_SUFFIX.sub("", stem) return stem
[docs] def sample_identity(path: Any) -> Dict[str, str]: """Parse plate/well/field/object identities from a crop filename. Unknown levels are returned as empty strings rather than guessed. The object identity is the augmentation-normalized full stem. :param path: crop path or file name. Augmentation suffixes are removed and the stem is split on underscores to find the plate, well and field parts. """ family = augmentation_family(path) parts = family.split("_") field_index = next( ( index for index in range(len(parts) - 1, -1, -1) if re.fullmatch(r"(?i)f\d+", parts[index]) is not None ), None, ) if ( field_index is None and len(parts) >= 3 and re.fullmatch(r"\d+", parts[-2]) is not None and re.fullmatch(r"(?i)(?:o)?\d+", parts[-1]) is not None ): field_index = len(parts) - 2 if field_index is not None and field_index >= 1: split_row_column = ( field_index >= 2 and re.fullmatch( r"(?i)(?:r\d+|[a-z])", parts[field_index - 2] ) is not None and re.fullmatch( r"(?i)(?:c\d+|\d+)", parts[field_index - 1] ) is not None ) plate_end = field_index - 2 if split_row_column else field_index - 1 plate = "_".join(parts[:plate_end]) if plate_end > 0 else "" well = "_".join(parts[:field_index]) field_id = "_".join(parts[: field_index + 1]) else: plate = parts[0] if parts and parts[0] else "" well = "_".join(parts[:2]) if len(parts) >= 2 else "" field_id = "_".join(parts[:3]) if len(parts) >= 3 else "" return { "sample": str(path), "basename": os.path.basename(str(path)), "augmentation_family": family, "object": family, "plate": plate, "well": well, "field": field_id, }
def _identity_sets(paths: Iterable[Any]) -> Dict[str, set]: """Return non-empty identity values for a path collection.""" result = { "exact": set(), "augmentation_family": set(), "object": set(), "field": set(), "well": set(), "plate": set(), } for path in paths: identity = sample_identity(path) result["exact"].add(os.path.abspath(str(path))) for level in result: if level == "exact": continue value = identity[level] if value: result[level].add(value) return result def _content_sha256(path: Any) -> Tuple[str, str]: """Return ``(sha256, error)`` for one file without loading it into memory.""" try: candidate = Path(str(path)) if not candidate.is_file(): return "", f"{candidate}: file does not exist" digest = hashlib.sha256() with candidate.open("rb") as handle: for chunk in iter(lambda: handle.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest(), "" except OSError as exc: return "", f"{path}: {type(exc).__name__}: {exc}" def _identity_sets_with_hashes( paths: Iterable[Any], *, hash_content: bool, ) -> Tuple[Dict[str, set], List[str]]: """Return identity sets, optionally including byte-content hashes.""" result = _identity_sets(paths) result["content_sha256"] = set() errors: List[str] = [] if hash_content: for path in paths: digest, error = _content_sha256(path) if digest: result["content_sha256"].add(digest) else: errors.append(error) return result, errors
[docs] def audit_split_leakage( train_paths: Sequence[Any], validation_paths: Sequence[Any], *, group_by: str = "well", raise_on_leakage: bool = False, split_name: str = "", hash_content: bool = False, require_identity: bool = False, ) -> LeakageReport: """Detect related images crossing a train/validation boundary. Exact/object/augmentation-family overlap is always critical. The requested ``group_by`` level is also critical: a well-grouped split permits the same plate on both sides but never the same well. :param train_paths: source paths used to fit the model. :param validation_paths: paths used only for evaluation. :param group_by: ``cell``, ``field``, ``well``, or ``plate``. Legacy ``none`` aliases ``cell``. :param raise_on_leakage: raise :class:`LeakageError` on a critical overlap. :param split_name: optional fold/split label stored in the report. :param hash_content: also compare file CONTENT, so a byte-identical copy under a different name is caught. Costs one read per file. :param require_identity: treat UNVERIFIABLE as critical. A filename that does not encode the requested identity, or a file that cannot be hashed, is otherwise only a warning -- so leaving this False means a clean report can still hide leakage nobody could check for. :returns: :class:`LeakageReport`. """ group_by = normalize_split_level(group_by) train, train_hash_errors = _identity_sets_with_hashes( train_paths, hash_content=hash_content, ) validation, validation_hash_errors = _identity_sets_with_hashes( validation_paths, hash_content=hash_content, ) overlap = { level: sorted(train[level] & validation[level]) for level in train } critical_candidates = [ "exact", "content_sha256", "augmentation_family", "object", ] if group_by != "cell": critical_candidates.append(group_by) critical = [ level for level in dict.fromkeys(critical_candidates) if overlap[level] ] warnings = [] if group_by == "cell": warnings.append( "The split is not grouped; shared well/field acquisition context " "can inflate performance even when exact objects do not overlap." ) missing_train = sum( not sample_identity(path)[group_by] for path in train_paths ) if group_by != "cell" else 0 missing_val = sum( not sample_identity(path)[group_by] for path in validation_paths ) if group_by != "cell" else 0 unverifiable = {} if missing_train or missing_val: unverifiable[group_by] = int(missing_train + missing_val) warnings.append( f"{missing_train} training and {missing_val} validation filename(s) " f"do not encode the requested {group_by} identity." ) if require_identity: critical.append(f"unverifiable_{group_by}") hash_errors = train_hash_errors + validation_hash_errors if hash_content and hash_errors: warnings.append( f"{len(hash_errors)} file(s) could not be content-hashed, so renamed " "byte-identical copies cannot be excluded for those samples." ) if require_identity: critical.append("unverifiable_content") report = LeakageReport( group_by=group_by, train_samples=len(train_paths), validation_samples=len(validation_paths), overlap_counts={level: len(values) for level, values in overlap.items()}, examples={level: values[:10] for level, values in overlap.items()}, split_name=str(split_name), critical_levels=critical, warnings=warnings, unverifiable_counts=unverifiable, hash_errors=hash_errors[:20], ) if raise_on_leakage and not report.passed: details = ", ".join( f"{level}={report.overlap_counts.get(level, report.unverifiable_counts.get(group_by, 1))}" for level in report.critical_levels ) raise LeakageError( f"Train/validation leakage detected ({details}). Rebuild folds " f"with cv_group_by={group_by!r}, split before augmentation, and " "preserve spaCR crop identities in filenames." ) return report
@dataclass
[docs] class FoldLeakageAudit: """Whole-CV proof that each related sample family belongs to one fold. :param group_by: canonical identity level required to remain within one validation fold. :param n_samples: number of source paths whose fold membership was audited. :param n_folds: number of train/validation fold pairs inspected. :param validation_membership_missing: up to twenty sample indexes that were never held out for validation. :param validation_membership_duplicate: up to twenty sample indexes held out in more than one fold. :param overlap_counts: identities assigned to multiple validation folds, counted at each exact, content, family, object, and acquisition level. :param examples: up to ten conflicting identities per level, annotated with the folds that contain them. :param critical_levels: completeness, overlap, identity, or label failures that make :attr:`passed` false. :param warnings: non-fatal caveats and explanations accompanying failures. :param unverifiable_counts: samples whose requested identity or optional byte-content hash could not be checked, counted by level. :param hash_errors: up to twenty file-specific content-hashing failures. :param split_name: stable label for this whole-CV audit record. """ group_by: str n_samples: int n_folds: int validation_membership_missing: List[int] = field(default_factory=list) validation_membership_duplicate: List[int] = field(default_factory=list) overlap_counts: Dict[str, int] = field(default_factory=dict) examples: Dict[str, List[str]] = field(default_factory=dict) critical_levels: List[str] = field(default_factory=list) warnings: List[str] = field(default_factory=list) unverifiable_counts: Dict[str, int] = field(default_factory=dict) hash_errors: List[str] = field(default_factory=list) split_name: str = "all_cv_folds" @property
[docs] def passed(self) -> bool: """Return True only when the fold partition is complete and isolated.""" return not self.critical_levels
[docs] def to_dict(self) -> Dict[str, Any]: """Return a JSON-serializable audit record.""" result = asdict(self) result["passed"] = self.passed return result
[docs] def audit_cv_folds( paths: Sequence[Any], folds: Sequence[Tuple[Sequence[int], Sequence[int]]], *, labels: Optional[Sequence[Any]] = None, group_by: str = "well", hash_content: bool = False, require_identity: bool = True, raise_on_leakage: bool = False, ) -> FoldLeakageAudit: """Verify partition coverage and identity isolation across every CV fold. This checks the fold assignment as one object, rather than trusting a sample of pairwise boundaries. Each index must be validation exactly once; exact paths, byte-identical content, source/augmentation families and the requested plate/well/field group must map to one held-out fold only. :param paths: one source path per sample, indexed by the fold indices. :param folds: ``(train_indices, validation_indices)`` per fold. Indices outside ``range(len(paths))`` are reported rather than ignored. :param labels: optional class per path, one value per path. Only used to describe the folds; it does not affect leakage detection. :param group_by: identity level that may not cross a fold boundary -- ``none``, ``field``, ``well`` or ``plate``. A well-grouped split permits the same plate on both sides but never the same well. :param hash_content: also compare file CONTENT, so a byte-identical copy under a different name is caught. Costs one read per file. :param require_identity: treat UNVERIFIABLE as critical. A filename that does not encode the requested identity, or a file that cannot be hashed, is otherwise only a warning -- so leaving this False means a clean report can still hide leakage nobody could check for. :param raise_on_leakage: raise :class:`LeakageError` instead of returning a report whose ``passed`` is False. :returns: a :class:`FoldLeakageAudit` carrying the per-level overlap counts, examples, and which levels were critical. :raises ValueError: for an unsupported ``group_by``, or ``labels`` whose length does not match ``paths``. """ group_by = normalize_split_level(group_by) n_samples = len(paths) membership: List[List[int]] = [[] for _ in range(n_samples)] warnings: List[str] = [] for fold_index, (train_indices, validation_indices) in enumerate(folds, 1): train_set = {int(index) for index in train_indices} validation_set = {int(index) for index in validation_indices} invalid = sorted( index for index in train_set | validation_set if index < 0 or index >= n_samples ) if invalid: raise ValueError( f"fold {fold_index} contains out-of-range indexes {invalid[:10]}" ) if train_set & validation_set: raise LeakageError( f"fold {fold_index} puts indexes in both train and validation: " f"{sorted(train_set & validation_set)[:10]}" ) for index in validation_set: membership[index].append(fold_index) missing = [index for index, owners in enumerate(membership) if not owners] duplicate = [index for index, owners in enumerate(membership) if len(owners) > 1] levels = ( "exact", "content_sha256", "augmentation_family", "object", "field", "well", "plate", ) owners_by_identity: Dict[str, Dict[str, set]] = { level: {} for level in levels } missing_identity = 0 hash_errors: List[str] = [] label_by_identity: Dict[str, Dict[str, set]] = { level: {} for level in ("content_sha256", "augmentation_family", "object") } label_values = list(labels) if labels is not None else None if label_values is not None and len(label_values) != n_samples: raise ValueError("labels must have one value per path") for index, path in enumerate(paths): identity = sample_identity(path) values = { "exact": os.path.abspath(str(path)), **{level: identity[level] for level in ( "augmentation_family", "object", "field", "well", "plate", )}, } digest = "" if hash_content: digest, error = _content_sha256(path) if error: hash_errors.append(error) values["content_sha256"] = digest if group_by != "cell" and not values[group_by]: missing_identity += 1 for level, value in values.items(): if not value: continue owners_by_identity[level].setdefault(value, set()).update( membership[index] ) if label_values is not None and level in label_by_identity: label_by_identity[level].setdefault(value, set()).add( str(label_values[index]) ) overlaps = { level: { value: sorted(owners) for value, owners in owners_by_identity[level].items() if len(owners) > 1 } for level in levels } critical = [] if missing: critical.append("validation_membership_missing") if duplicate: critical.append("validation_membership_duplicate") protected = ["exact", "content_sha256", "augmentation_family", "object"] if group_by != "cell": protected.append(group_by) critical.extend(level for level in protected if overlaps[level]) conflicts = { level: sorted( value for value, assigned in assignments.items() if len(assigned) > 1 ) for level, assignments in label_by_identity.items() } if any(conflicts.values()): critical.append("conflicting_labels") warnings.append( "Related crops carry different class labels; fix annotations before " "training even when all copies happen to be in one fold." ) unverifiable = {} if missing_identity: unverifiable[group_by] = missing_identity warnings.append( f"{missing_identity} sample(s) do not encode {group_by} identity." ) if require_identity: critical.append(f"unverifiable_{group_by}") if hash_content and hash_errors: unverifiable["content_sha256"] = len(hash_errors) warnings.append(f"{len(hash_errors)} sample(s) could not be hashed.") if require_identity: critical.append("unverifiable_content") examples = { level: [ f"{value} -> folds {','.join(map(str, owners))}" for value, owners in list(overlaps[level].items())[:10] ] for level in levels } examples["conflicting_labels"] = [ f"{level}:{value}" for level, values in conflicts.items() for value in values[:10] ][:10] audit = FoldLeakageAudit( group_by=group_by, n_samples=n_samples, n_folds=len(folds), validation_membership_missing=missing[:20], validation_membership_duplicate=duplicate[:20], overlap_counts={ level: len(values) for level, values in overlaps.items() }, examples=examples, critical_levels=list(dict.fromkeys(critical)), warnings=warnings, unverifiable_counts=unverifiable, hash_errors=hash_errors[:20], ) if raise_on_leakage and not audit.passed: raise LeakageError( "Cross-validation leakage audit failed: " + ", ".join(audit.critical_levels) ) return audit
[docs] def dataset_split_paths(root: Any, split: str) -> List[str]: """Return sorted image paths under ``root/<split>/<class>/``. :param root: dataset folder that contains the split folders; ``~`` is expanded. :param split: name of the split subfolder, such as ``train`` or ``test``. A missing folder returns an empty list; files with an image or ``.npy`` suffix are collected recursively below it. """ folder = Path(str(root)).expanduser() / str(split) if not folder.is_dir(): return [] suffixes = {".png", ".jpg", ".jpeg", ".tif", ".tiff", ".bmp", ".npy"} return sorted( str(path) for path in folder.rglob("*") if path.is_file() and path.suffix.lower() in suffixes )
[docs] def audit_dataset_splits( root: Any, *, group_by: str = "well", hash_content: bool = True, require_identity: bool = True, raise_on_leakage: bool = False, ) -> LeakageReport: """Audit the permanent ``train/`` versus ``test/`` dataset boundary. :param root: dataset directory holding ``train/`` and ``test/``. Both are searched recursively, and only ``.png``, ``.jpg``, ``.jpeg``, ``.tif``, ``.tiff``, ``.bmp`` and ``.npy`` files are collected, so a folder of any other format audits as empty. :param group_by: identity level that may not appear on both sides -- ``none``, ``field``, ``well`` or ``plate``. The default ``well`` still permits the same plate in train and test. It is validated only after the images are found, so a bad value over an empty tree reports the missing images instead. :param hash_content: defaults to True here, unlike :func:`audit_split_leakage`, so a byte-identical copy saved under a different name fails the audit -- at the cost of one read per file. :param require_identity: defaults to True here: filenames that do not encode the ``group_by`` level, and files that cannot be hashed, become critical instead of a warning, so an unverifiable split cannot report as clean. :param raise_on_leakage: raise :class:`LeakageError` instead of returning a report whose ``passed`` is False. :raises FileNotFoundError: when either side collects no image -- a missing or unreadable folder is never treated as a passing split. :returns: a :class:`LeakageReport` whose ``split_name`` is always ``train_vs_test``; the ``test/`` side is counted as ``validation_samples``. """ train_paths = dataset_split_paths(root, "train") test_paths = dataset_split_paths(root, "test") if not train_paths or not test_paths: missing = [ name for name, values in (("train", train_paths), ("test", test_paths)) if not values ] raise FileNotFoundError( f"Cannot audit dataset leakage: no images found in {', '.join(missing)} " f"under {Path(str(root)).expanduser()}." ) return audit_split_leakage( train_paths, test_paths, group_by=group_by, raise_on_leakage=raise_on_leakage, split_name="train_vs_test", hash_content=hash_content, require_identity=require_identity, )
[docs] def write_leakage_audit(path: Any, audit: Any) -> Path: """Atomically write a leakage report/audit as JSON and return its path. :param path: destination JSON file; ``~`` is expanded and missing parent folders are created. It is replaced atomically. :param audit: leakage report or audit to write: any object with a ``to_dict()`` method, or a mapping. """ destination = Path(str(path)).expanduser() destination.parent.mkdir(parents=True, exist_ok=True) payload = audit.to_dict() if hasattr(audit, "to_dict") else dict(audit) temporary = destination.with_name(destination.name + ".tmp") temporary.write_text( json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8" ) os.replace(temporary, destination) return destination
[docs] def normalize_probabilities( probabilities: Any, *, n_classes: Optional[int] = None, ) -> np.ndarray: """Return a finite, row-normalized ``(n_samples, n_classes)`` matrix. :param probabilities: a one-dimensional positive-class vector, a one-column matrix (both expanded to two classes), or a samples-by-classes matrix. Values must be finite and within 0 to 1, and no row may sum to zero. """ matrix = np.asarray(probabilities, dtype=float) if matrix.ndim == 1: matrix = np.column_stack([1.0 - matrix, matrix]) elif matrix.ndim == 2 and matrix.shape[1] == 1: values = matrix[:, 0] matrix = np.column_stack([1.0 - values, values]) if matrix.ndim != 2 or matrix.shape[1] < 2: raise ValueError( "probabilities must be a one-dimensional positive-class vector " "or a two-dimensional matrix with at least two classes." ) if n_classes is not None and matrix.shape[1] != int(n_classes): raise ValueError( f"Probability matrix has {matrix.shape[1]} columns but " f"{n_classes} classes were declared." ) if not np.isfinite(matrix).all(): raise ValueError("probabilities contain NaN or infinite values.") if (matrix < 0).any() or (matrix > 1).any(): raise ValueError("probabilities must lie between 0 and 1.") totals = matrix.sum(axis=1, keepdims=True) if (totals <= 0).any(): raise ValueError("At least one probability row sums to zero.") return matrix / totals
[docs] def expected_calibration_error( y_true: Sequence[int], probabilities: Any, *, n_bins: int = 10, ) -> float: """Return top-label expected calibration error (ECE). :param y_true: true class index per sample, aligned with the rows of ``probabilities``. :param probabilities: predicted probabilities, either a positive-class vector or a samples-by-classes matrix; validated and row-normalised by :func:`normalize_probabilities`. Its length must equal that of ``y_true``. """ y = np.asarray(y_true, dtype=int) probs = normalize_probabilities(probabilities) if len(y) != len(probs): raise ValueError("y_true and probabilities must have equal length.") if len(y) == 0: return float("nan") n_bins = max(2, int(n_bins)) confidence = probs.max(axis=1) predicted = probs.argmax(axis=1) correct = predicted == y edges = np.linspace(0.0, 1.0, n_bins + 1) total = float(len(y)) ece = 0.0 for index in range(n_bins): if index == n_bins - 1: mask = (confidence >= edges[index]) & (confidence <= edges[index + 1]) else: mask = (confidence >= edges[index]) & (confidence < edges[index + 1]) if not mask.any(): continue ece += (mask.sum() / total) * abs( float(correct[mask].mean()) - float(confidence[mask].mean()) ) return float(ece)
def _temperature_probabilities(probabilities: np.ndarray, temperature: float) -> np.ndarray: """Apply temperature scaling to an already normalized probability matrix.""" temperature = max(float(temperature), 1e-6) logits = np.log(np.clip(probabilities, 1e-12, 1.0)) / temperature logits -= logits.max(axis=1, keepdims=True) exp = np.exp(logits) return exp / exp.sum(axis=1, keepdims=True)
[docs] def fit_temperature( y_true: Sequence[int], probabilities: Any, ) -> float: """Fit one scalar temperature by minimizing multiclass log loss. :param y_true: true class index per sample; at least two samples and two distinct classes are required. :param probabilities: uncalibrated predicted probabilities, a positive-class vector or a samples-by-classes matrix, normalised by :func:`normalize_probabilities` before the fit. """ from scipy.optimize import minimize_scalar y = np.asarray(y_true, dtype=int) probs = normalize_probabilities(probabilities) if len(y) < 2 or np.unique(y).size < 2: raise ValueError( "Temperature calibration needs at least two classes and two samples." ) def objective(log_temperature: float) -> float: """Return multiclass log loss after applying an exponentiated scale.""" calibrated = _temperature_probabilities( probs, math.exp(float(log_temperature)), ) return float(log_loss(y, calibrated, labels=np.arange(probs.shape[1]))) result = minimize_scalar( objective, bounds=(math.log(0.05), math.log(20.0)), method="bounded", ) if not result.success: raise RuntimeError(f"Temperature fitting failed: {result.message}") return float(math.exp(float(result.x)))
[docs] def cross_calibrate_probabilities( y_true: Sequence[int], probabilities: Any, fold_ids: Sequence[Any], *, method: str = "temperature", warnings_out: Optional[List[str]] = None, ) -> Tuple[np.ndarray, Dict[str, float]]: """Calibrate each held-out fold using every *other* out-of-fold prediction. This cross-fitting prevents a sample from fitting the calibrator that is evaluated on that same sample. :param y_true: class INDEX per sample, read only to fit the temperatures. A value outside the probability columns is REFUSED -- it used to make every fold's fit fail, fall back to temperature 1.0, and return uncalibrated probabilities while reporting that calibration ran. :param probabilities: predicted probabilities; a 1-D positive-class vector is expanded to two columns and every row is renormalized, so the result is always a new normalized matrix -- never the object passed in, even when no calibration is applied. :param fold_ids: held-out block per sample; any hashable value works, and each distinct value is calibrated from all the others, so at least two distinct values are required. The returned map is keyed by ``str(fold_id)``. :param method: ``temperature``, or ``none``/``off``/``false`` (also ``None`` and the empty string) to return the normalized probabilities unchanged with an empty temperature map. Case and surrounding whitespace are ignored; anything else is refused rather than skipped. :param warnings_out: list appended in place when a fold cannot be calibrated from the others. The message is printed either way, so leaving this None only discards the machine-readable copy. :raises ValueError: on mismatched lengths, an unrecognized ``method``, or fewer than two distinct ``fold_ids`` when calibrating. :returns: ``(calibrated_probabilities, temperature_by_held_out_fold)``. """ normalized_method = str(method or "none").strip().lower() probs = normalize_probabilities(probabilities) y = np.asarray(y_true, dtype=int) folds = np.asarray(fold_ids) if len(y) and probs.shape[1]: stray = sorted({int(v) for v in y if v < 0 or v >= probs.shape[1]}) if stray: raise ValueError( f"y_true contains class indices {stray} outside the " f"{probs.shape[1]} probability columns; calibration would " "silently do nothing.") if len(y) != len(probs) or len(folds) != len(y): raise ValueError( "y_true, probabilities, and fold_ids must have equal length." ) if normalized_method in {"none", "off", "false"}: return probs.copy(), {} if normalized_method != "temperature": raise ValueError( "calibration_method must be 'none' or 'temperature', not " f"{method!r}." ) unique_folds = list(pd.unique(folds)) if len(unique_folds) < 2: raise ValueError( "Cross-fitted calibration requires predictions from at least two folds." ) calibrated = np.empty_like(probs) temperatures: Dict[str, float] = {} for fold in unique_folds: target = folds == fold fit = ~target try: temperature = fit_temperature(y[fit], probs[fit]) except (ValueError, RuntimeError) as exc: temperature = 1.0 warning = ( f"Held-out fold {fold!r} could not be temperature-calibrated " f"from the other folds ({type(exc).__name__}: {exc}); its raw " "probabilities were retained." ) print(f"Warning: classifier calibration: {warning}") if warnings_out is not None: warnings_out.append(warning) calibrated[target] = _temperature_probabilities( probs[target], temperature, ) temperatures[str(fold)] = temperature return calibrated, temperatures
#: Columns of the frame :func:`calibration_table` returns, named here so an #: empty result still has them. Without this a table with no rows is shape #: (0, 0) and every downstream column access raises KeyError. CALIBRATION_COLUMNS = ( "class_index", "class_name", "bin", "bin_lower", "bin_upper", "n", "mean_confidence", "observed_frequency", "calibration_gap", )
[docs] def calibration_table( y_true: Sequence[int], probabilities: Any, *, classes: Optional[Sequence[str]] = None, n_bins: int = 10, ) -> pd.DataFrame: """Return per-class reliability bins for calibration plots. :param y_true: class INDEX per sample, compared column by column. A label outside the probability columns is REFUSED: it matches no class, so it used to read ``observed_frequency`` 0.0 everywhere and render as a catastrophically miscalibrated curve rather than an error. :param probabilities: predicted probabilities; a 1-D positive-class vector is expanded to two columns and every row is renormalized. :param classes: display names in column order. An EMPTY sequence falls back to ``class_0 ... class_n``; a non-empty one whose length disagrees with the columns is refused. :param n_bins: equal-width confidence bins over ``[0, 1]``, truncated to an int and floored at 2, so 0, 1 and any negative value all give two bins. The top bin is closed on the right, so confidence 1.0 lands in it rather than falling out of the table. :raises ValueError: when ``y_true`` and ``probabilities`` differ in length, ``classes`` has the wrong length, or a class index falls outside the probability columns. :returns: one row per class and NON-EMPTY bin, so the frame is shorter than ``n_classes * n_bins`` rows. An empty result still carries :data:`CALIBRATION_COLUMNS`, so a column can be indexed on it. """ y = np.asarray(y_true, dtype=int) probs = normalize_probabilities(probabilities) if len(y) != len(probs): raise ValueError("y_true and probabilities must have equal length.") n_classes = probs.shape[1] if len(y) and n_classes: stray = sorted({int(v) for v in y if v < 0 or v >= n_classes}) if stray: raise ValueError( f"y_true contains class indices {stray} outside the " f"{n_classes} probability columns.") names = list(classes or [f"class_{i}" for i in range(n_classes)]) if len(names) != n_classes: raise ValueError("classes and probability columns must have equal length.") edges = np.linspace(0.0, 1.0, max(2, int(n_bins)) + 1) rows = [] for class_index, class_name in enumerate(names): observed = y == class_index confidence = probs[:, class_index] for bin_index in range(len(edges) - 1): if bin_index == len(edges) - 2: mask = ( (confidence >= edges[bin_index]) & (confidence <= edges[bin_index + 1]) ) else: mask = ( (confidence >= edges[bin_index]) & (confidence < edges[bin_index + 1]) ) if not mask.any(): continue rows.append({ "class_index": class_index, "class_name": class_name, "bin": bin_index + 1, "bin_lower": float(edges[bin_index]), "bin_upper": float(edges[bin_index + 1]), "n": int(mask.sum()), "mean_confidence": float(confidence[mask].mean()), "observed_frequency": float(observed[mask].mean()), "calibration_gap": float( observed[mask].mean() - confidence[mask].mean() ), }) return pd.DataFrame(rows, columns=CALIBRATION_COLUMNS)
def _metric_summary( y_true: np.ndarray, probabilities: np.ndarray, *, n_bins: int, ) -> Dict[str, Any]: """Compute scalar classifier metrics for one sample group.""" predicted = probabilities.argmax(axis=1) class_indices = np.arange(probabilities.shape[1]) one_hot = np.eye(probabilities.shape[1], dtype=float)[y_true] per_class_recall = recall_score( y_true, predicted, labels=class_indices, average=None, zero_division=0, ) supported = np.bincount( y_true, minlength=probabilities.shape[1], ) > 0 return { "n": int(len(y_true)), "accuracy": float(accuracy_score(y_true, predicted)), "balanced_accuracy": float(per_class_recall[supported].mean()), "f1_macro": float(f1_score( y_true, predicted, average="macro", zero_division=0, )), "f1_weighted": float(f1_score( y_true, predicted, average="weighted", zero_division=0, )), "precision_macro": float(precision_score( y_true, predicted, average="macro", zero_division=0, )), "recall_macro": float(recall_score( y_true, predicted, average="macro", zero_division=0, )), "log_loss": float(log_loss( y_true, probabilities, labels=class_indices, )), "brier_multiclass": float(np.mean(np.sum( (probabilities - one_hot) ** 2, axis=1, ))), "expected_calibration_error": expected_calibration_error( y_true, probabilities, n_bins=n_bins, ), "mean_confidence": float(probabilities.max(axis=1).mean()), }
[docs] def evaluate_predictions( y_true: Sequence[int], probabilities: Any, sample_paths: Sequence[Any], *, classes: Optional[Sequence[str]] = None, fold_ids: Optional[Sequence[Any]] = None, calibration_method: str = "none", calibration_bins: int = 10, ) -> Dict[str, Any]: """Build overall, confusion, calibration, and per-plate evaluation tables. :param y_true: true class INDEX per sample, in ``range(n_classes)``. A value outside that range is refused rather than clipped. :param probabilities: ``(n_samples, n_classes)`` predicted probabilities. Its column count defines ``n_classes``. :param sample_paths: one source path per sample, used to derive the per-plate tables. Must be the same length as ``y_true``. :param classes: display names for the columns, in column order. Defaults to ``class_0 ... class_n``; a length that disagrees with the probability columns is refused. :param fold_ids: which held-out fold each sample came from. REQUIRED for temperature calibration and unused otherwise -- the temperature is fitted per held-out fold so no sample is calibrated on itself, which needs at least two distinct folds. :param calibration_method: ``'none'`` (default) or ``'temperature'``. Anything else is refused rather than silently ignored. :param calibration_bins: bin count for the expected-calibration-error estimate. More bins resolve the reliability curve better and make each bin noisier. :raises ValueError: on mismatched lengths, a class index outside the probability columns, an unknown ``calibration_method``, or temperature calibration with fewer than two distinct folds. :returns: the evaluation bundle, a dict with six keys: * ``summary`` — scalar metrics (``n``, accuracy, balanced accuracy, macro/weighted F1, macro precision/recall, log loss, multiclass Brier, expected calibration error, mean confidence) plus ``classes``, ``n_classes``, ``calibration_method``, ``raw_expected_calibration_error``, ``temperatures_by_held_out_fold``, ``calibration_warnings`` and ``probability_column_names``; * ``predictions`` — one row per sample with ``fold``, the :func:`sample_identity` columns, ``true_label`` / ``true_class``, ``predicted_label`` / ``predicted_class``, ``correct``, ``confidence`` (the calibrated probability of the chosen class) and a ``raw_prob_<name>`` / ``prob_<name>`` pair per class, where ``<name>`` is the sanitized, de-duplicated class name listed in ``probability_column_names`` rather than the class name itself; * ``confusion_counts`` — counts indexed and columned by class name; * ``confusion_normalized`` — the same matrix divided by its true-class row totals, with all-zero rows left at zero; * ``per_plate`` — the same scalar metrics per ``plate`` group, with ``plate`` as the first column; * ``calibration`` — the :func:`calibration_table` reliability bins for the calibrated probabilities. """ y = np.asarray(y_true, dtype=int) raw = normalize_probabilities(probabilities) n_classes = raw.shape[1] names = list(classes or [f"class_{i}" for i in range(n_classes)]) calibration_bins = int(calibration_bins) if calibration_bins < 2: raise ValueError("calibration_bins must be at least 2.") if len(y) != len(raw) or len(sample_paths) != len(y): raise ValueError( "y_true, probabilities, and sample_paths must have equal length." ) if len(names) != n_classes: raise ValueError("classes and probability columns must have equal length.") if len(set(names)) != len(names): raise ValueError("class names must be unique.") if len(y) == 0: raise ValueError("At least one prediction is required.") if (y < 0).any() or (y >= n_classes).any(): raise ValueError("y_true contains a label outside the class schema.") folds = ( np.asarray(fold_ids) if fold_ids is not None else np.zeros(len(y), dtype=int) ) if len(folds) != len(y): raise ValueError( "fold_ids must have the same length as y_true and probabilities." ) calibration_warnings: List[str] = [] calibrated, temperatures = cross_calibrate_probabilities( y, raw, folds, method=calibration_method, warnings_out=calibration_warnings, ) if str(calibration_method).lower() not in {"none", "off", "false"} else ( raw.copy(), {} ) identities = pd.DataFrame([sample_identity(path) for path in sample_paths]) predicted = calibrated.argmax(axis=1) frame = identities.copy() frame.insert(0, "fold", folds) frame["true_label"] = y frame["true_class"] = [names[index] for index in y] frame["predicted_label"] = predicted frame["predicted_class"] = [names[index] for index in predicted] frame["correct"] = predicted == y frame["confidence"] = calibrated.max(axis=1) safe_names = [] used_safe_names = set() for index, name in enumerate(names): base = ( re.sub(r"[^A-Za-z0-9_.-]+", "_", str(name)).strip("_") or f"class_{index}" ) safe_name = base suffix = 2 while safe_name in used_safe_names: safe_name = f"{base}_{suffix}" suffix += 1 used_safe_names.add(safe_name) safe_names.append(safe_name) frame[f"raw_prob_{safe_name}"] = raw[:, index] frame[f"prob_{safe_name}"] = calibrated[:, index] summary = _metric_summary(y, calibrated, n_bins=calibration_bins) summary.update({ "classes": names, "n_classes": n_classes, "calibration_method": str(calibration_method or "none").lower(), "raw_expected_calibration_error": expected_calibration_error( y, raw, n_bins=calibration_bins, ), "temperatures_by_held_out_fold": temperatures, "calibration_warnings": calibration_warnings, "probability_column_names": safe_names, }) counts = confusion_matrix(y, predicted, labels=np.arange(n_classes)) row_totals = counts.sum(axis=1, keepdims=True) normalized = np.divide( counts, row_totals, out=np.zeros_like(counts, dtype=float), where=row_totals != 0, ) confusion_counts = pd.DataFrame(counts, index=names, columns=names) confusion_normalized = pd.DataFrame( normalized, index=names, columns=names, ) per_plate_rows = [] for plate, group in frame.groupby("plate", dropna=False): indices = group.index.to_numpy(dtype=int) metrics = _metric_summary( y[indices], calibrated[indices], n_bins=calibration_bins, ) metrics["plate"] = plate or "unknown" per_plate_rows.append(metrics) per_plate = pd.DataFrame(per_plate_rows) columns = ["plate", *[c for c in per_plate if c != "plate"]] per_plate = per_plate[columns] return { "summary": summary, "predictions": frame, "confusion_counts": confusion_counts, "confusion_normalized": confusion_normalized, "per_plate": per_plate, "calibration": calibration_table( y, calibrated, classes=names, n_bins=calibration_bins, ), }
[docs] def nested_group_folds( labels: Sequence[int], *, outer_splits: int, inner_splits: int, groups: Optional[Sequence[Any]] = None, seed: int = 0, ) -> List[Dict[str, Any]]: """Build nested stratified/grouped outer and inner index partitions. Inner indexes are returned in the original/global coordinate system. :param labels: class label per sample, coerced to int. Only its length and class balance matter -- the folds carry indices, not data. :param outer_splits: outer fold count, coerced with ``int`` (``2.9`` gives 2); below 2 is refused. :param inner_splits: inner fold count built inside each outer TRAINING set, so it is bounded by that subset, not by the dataset: three inner folds over four samples fails on the outer training half even though the outer split itself succeeded. :param groups: optional group key per sample (well or plate id) kept whole within a fold at BOTH levels. None gives a plain stratified split that will scatter crops of the same well across folds. :param seed: outer folds use ``seed``; the inner folds of outer fold ``k`` use ``seed + k``, so inner partitions differ between outer folds instead of repeating one layout. Must be an int -- a float reaches ``numpy.random.default_rng`` and raises ``TypeError``. :raises ValueError: when either split count is below 2, or when the sample count (or the number of distinct groups) cannot supply that many folds. :returns: one dict per outer fold, with ``outer_fold`` (1-based), ``train``, ``validation`` and ``inner`` -- a list of ``(train, validation)`` index pairs expressed in GLOBAL indices, not as positions inside ``train``. """ from .io import make_cv_folds outer_splits = int(outer_splits) inner_splits = int(inner_splits) if outer_splits < 2 or inner_splits < 2: raise ValueError("outer_splits and inner_splits must both be at least 2.") y = np.asarray(labels, dtype=int) group_values = None if groups is None else np.asarray(groups) outer = make_cv_folds( y, outer_splits, groups=group_values, seed=seed, ) result = [] for outer_index, (outer_train, outer_validation) in enumerate( outer, start=1, ): inner_groups = ( None if group_values is None else group_values[outer_train] ) relative_inner = make_cv_folds( y[outer_train], inner_splits, groups=inner_groups, seed=seed + outer_index, ) inner = [ (outer_train[train_relative], outer_train[val_relative]) for train_relative, val_relative in relative_inner ] result.append({ "outer_fold": outer_index, "train": outer_train, "validation": outer_validation, "inner": inner, }) return result
[docs] def write_evaluation_bundle( output_dir: Any, evaluation: Mapping[str, Any], *, leakage_reports: Optional[Sequence[LeakageReport]] = None, ) -> Path: """Atomically write a complete evaluation bundle and diagnostic figures. :param output_dir: folder that receives the bundle; created with its parents if needed. :param evaluation: mapping as returned by :func:`evaluate_predictions`, with ``summary``, ``predictions``, ``confusion_counts``, ``confusion_normalized``, ``per_plate`` and ``calibration`` entries. """ destination = Path(output_dir) destination.mkdir(parents=True, exist_ok=True) def write_json(name: str, payload: Any) -> None: """Atomically replace ``name`` with stable indented JSON.""" path = destination / name temporary = path.with_name(f".{path.name}.tmp") temporary.write_text( json.dumps(payload, indent=2, sort_keys=True, allow_nan=True), encoding="utf-8", ) temporary.replace(path) def write_csv(name: str, frame: pd.DataFrame, *, index: bool = False) -> None: """Atomically replace ``name`` with ``frame`` and optional index.""" path = destination / name temporary = path.with_name(f".{path.name}.tmp") frame.to_csv(temporary, index=index) temporary.replace(path) write_json(EVALUATION_FILES["summary"], evaluation["summary"]) write_csv(EVALUATION_FILES["predictions"], evaluation["predictions"]) write_csv( EVALUATION_FILES["confusion_counts"], evaluation["confusion_counts"], index=True, ) write_csv( EVALUATION_FILES["confusion_normalized"], evaluation["confusion_normalized"], index=True, ) write_csv(EVALUATION_FILES["per_plate"], evaluation["per_plate"]) write_csv(EVALUATION_FILES["calibration"], evaluation["calibration"]) reports = [report.to_dict() for report in (leakage_reports or [])] write_json(EVALUATION_FILES["leakage"], { "passed": all(report["passed"] for report in reports), "folds": reports, }) figure_warnings = [] try: _write_confusion_figure( evaluation["confusion_normalized"], destination / EVALUATION_FILES["confusion_figure"], ) except Exception as exc: figure_warnings.append( f"Confusion figure failed ({type(exc).__name__}: {exc})." ) try: _write_calibration_figure( evaluation["calibration"], destination / EVALUATION_FILES["calibration_figure"], ) except Exception as exc: figure_warnings.append( f"Calibration figure failed ({type(exc).__name__}: {exc})." ) manifest = { "schema_version": 1, "files": EVALUATION_FILES, "summary": evaluation["summary"], "leakage_passed": all(report["passed"] for report in reports), "warnings": figure_warnings, } write_json(EVALUATION_FILES["manifest"], manifest) for warning in figure_warnings: print(f"Warning: classifier evaluation: {warning}") return destination / EVALUATION_FILES["manifest"]
def _write_confusion_figure(frame: pd.DataFrame, path: Path) -> None: """Render a normalized confusion heatmap.""" import matplotlib.pyplot as plt with figure_style(theme_target()): fig, axis = plt.subplots(figsize=(6, 5)) from .figures.bundle import _register_figure_data _register_figure_data(fig, frame, kind="heatmap", matrix=True, title="Out-of-fold confusion matrix") image = axis.imshow(frame.to_numpy(dtype=float), vmin=0, vmax=1, cmap="Blues") axis.set_xticks(np.arange(len(frame.columns)), labels=frame.columns, rotation=45, ha="right") axis.set_yticks(np.arange(len(frame.index)), labels=frame.index) axis.set_xlabel("Predicted") axis.set_ylabel("True") axis.set_title("Out-of-fold confusion matrix") for row in range(len(frame.index)): for column in range(len(frame.columns)): value = float(frame.iloc[row, column]) axis.text(column, row, f"{value:.2f}", ha="center", va="center", color="white" if value > 0.5 else "black") fig.colorbar(image, ax=axis, label="Row-normalized fraction") fig.tight_layout() from .plot import save_figure save_figure(fig, path, fmt="png", close=True) def _write_calibration_figure(frame: pd.DataFrame, path: Path) -> None: """Render one reliability curve per class.""" import matplotlib.pyplot as plt with figure_style(theme_target()): fig, axis = plt.subplots(figsize=(6, 5)) from .figures.bundle import _register_figure_data _register_figure_data(fig, frame, x="mean_confidence", y="observed_frequency", hue="class_name", kind="line") axis.plot([0, 1], [0, 1], linestyle="--", color="#777777", label="Perfect calibration") for class_name, group in frame.groupby("class_name"): axis.plot( group["mean_confidence"], group["observed_frequency"], marker="o", label=str(class_name), ) axis.set_xlim(0, 1) axis.set_ylim(0, 1) axis.set_xlabel("Mean predicted probability") axis.set_ylabel("Observed frequency") axis.set_title("Out-of-fold calibration") axis.legend(loc="best") fig.tight_layout() from .plot import save_figure save_figure(fig, path, fmt="png", close=True)
[docs] def find_evaluation_bundles(root: Any) -> List[Path]: """Return evaluation manifests below ``root``, newest first. :param root: folder searched recursively for evaluation manifests, or a manifest file (returned alone) or any other file (its folder is searched). A missing path raises :class:`FileNotFoundError`. """ source = Path(root).expanduser() if source.is_file(): if source.name == EVALUATION_FILES["manifest"]: return [source] source = source.parent if not source.exists(): raise FileNotFoundError(f"Evaluation source does not exist: {source}") manifests = list(source.rglob(EVALUATION_FILES["manifest"])) return sorted( manifests, key=lambda path: path.stat().st_mtime, reverse=True, )
[docs] def load_evaluation_bundle(path: Any) -> Dict[str, Any]: """Load one evaluation bundle for the Qt workbench. :param path: an evaluation manifest file, or the bundle folder that contains it; ``~`` is expanded. Tables named by the manifest that are absent load as empty frames. """ source = Path(path).expanduser() manifest_path = ( source if source.is_file() else source / EVALUATION_FILES["manifest"] ) if not manifest_path.is_file(): raise FileNotFoundError( f"No {EVALUATION_FILES['manifest']} found at {source}." ) folder = manifest_path.parent manifest = json.loads(manifest_path.read_text(encoding="utf-8")) def read_csv(key: str, **kwargs) -> pd.DataFrame: """Load a manifest-named CSV, or an empty frame when it is absent.""" file_name = manifest.get("files", EVALUATION_FILES).get( key, EVALUATION_FILES[key], ) path_value = folder / file_name return pd.read_csv(path_value, **kwargs) if path_value.is_file() else pd.DataFrame() leakage_path = folder / manifest.get("files", EVALUATION_FILES).get( "leakage", EVALUATION_FILES["leakage"], ) leakage = ( json.loads(leakage_path.read_text(encoding="utf-8")) if leakage_path.is_file() else {} ) return { "path": manifest_path, "manifest": manifest, "summary": manifest.get("summary", {}), "predictions": read_csv("predictions"), "confusion_counts": read_csv("confusion_counts", index_col=0), "confusion_normalized": read_csv( "confusion_normalized", index_col=0, ), "per_plate": read_csv("per_plate"), "calibration": read_csv("calibration"), "leakage": leakage, }