Coverage for imcluster/evaluate.py: 100.00%
61 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-12 01:58 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-12 01:58 +0000
1"""Evaluation metrics for image clustering results."""
3from pathlib import Path
5import numpy as np
6import pandas as pd
7from numpy.typing import ArrayLike
8from rich.console import Console
9from rich.table import Table
10from scipy.optimize import linear_sum_assignment
11from sklearn.metrics import adjusted_rand_score, normalized_mutual_info_score
13from .io import ImclusterIO
15console = Console()
18def clustering_accuracy(expected: ArrayLike, predicted: ArrayLike) -> float:
19 """Return clustering accuracy under the optimal class-label assignment.
21 Args:
22 expected: Ground-truth class labels.
23 predicted: Predicted cluster labels.
25 Returns:
26 The fraction of samples assigned correctly after optimally matching
27 cluster IDs to expected classes.
29 Raises:
30 ValueError: If the label arrays are empty or have different lengths.
31 """
32 expected_labels = np.asarray(expected)
33 predicted_labels = np.asarray(predicted)
34 if not len(expected_labels) or len(expected_labels) != len(predicted_labels):
35 raise ValueError("expected and predicted labels must have equal nonzero length")
37 expected_values, expected_codes = np.unique(expected_labels, return_inverse=True)
38 predicted_values, predicted_codes = np.unique(predicted_labels, return_inverse=True)
39 contingency = np.zeros((len(expected_values), len(predicted_values)), dtype=int)
40 np.add.at(contingency, (expected_codes, predicted_codes), 1)
41 rows, columns = linear_sum_assignment(contingency, maximize=True)
42 return float(contingency[rows, columns].sum() / len(expected_labels))
45def evaluate_clustering(
46 imcluster_io: ImclusterIO,
47 expected_csv: str | Path,
48 cluster_column: str,
49) -> dict[str, float]:
50 """Evaluate cached cluster assignments against filename-based classes.
52 Args:
53 imcluster_io: Image collection containing filenames and cluster labels.
54 expected_csv: CSV with ``filename`` and ``class`` columns.
55 cluster_column: DataFrame column containing predicted cluster labels.
57 Returns:
58 NMI, ARI, and optimally matched clustering accuracy.
60 Raises:
61 ValueError: If the CSV or clustering results cannot be aligned.
62 """
63 expected = pd.read_csv(expected_csv, dtype=str)
64 required_columns = {"filename", "class"}
65 missing_columns = required_columns.difference(expected.columns)
66 if missing_columns:
67 missing = ", ".join(sorted(missing_columns))
68 raise ValueError(f"Expected classes CSV is missing columns: {missing}")
69 seen_filenames: set[str] = set()
70 duplicates: list[str] = []
71 for filename in expected["filename"].tolist():
72 if filename in seen_filenames and filename not in duplicates:
73 duplicates.append(filename)
74 seen_filenames.add(filename)
75 if duplicates:
76 raise ValueError(
77 "Expected classes CSV contains duplicate filenames: "
78 + ", ".join(duplicates)
79 )
80 if expected["class"].isna().any() or (expected["class"].str.strip() == "").any():
81 raise ValueError("Expected classes CSV contains an empty class")
82 if cluster_column not in imcluster_io.df:
83 raise ValueError(f"Missing clustering results column: {cluster_column}")
85 classes_by_filename = expected.set_index("filename")["class"]
86 filenames = imcluster_io.df["filenames"]
87 missing_filenames = sorted(set(filenames).difference(classes_by_filename.index))
88 if missing_filenames:
89 raise ValueError(
90 "Expected classes CSV has no class for: " + ", ".join(missing_filenames)
91 )
93 expected_labels = filenames.map(classes_by_filename).to_numpy()
94 predicted_labels = imcluster_io.df[cluster_column].to_numpy()
95 return {
96 "NMI": float(normalized_mutual_info_score(expected_labels, predicted_labels)),
97 "ARI": float(adjusted_rand_score(expected_labels, predicted_labels)),
98 "ACC": clustering_accuracy(expected_labels, predicted_labels),
99 }
102def print_evaluation(metrics: dict[str, float]) -> None:
103 """Print clustering metrics as a Rich table."""
104 descriptions = {
105 "NMI": "Normalized Mutual Information",
106 "ARI": "Adjusted Rand Index",
107 "ACC": "Clustering Accuracy",
108 }
109 table = Table(title="Clustering evaluation", show_header=True)
110 table.add_column("Metric", style="bold cyan")
111 table.add_column("Description")
112 table.add_column("Score", justify="right", style="green")
113 for metric, score in metrics.items():
114 table.add_row(metric, descriptions[metric], f"{score:.4f}")
115 console.print(table)
118def write_evaluation(metrics: dict[str, float], output_csv: str | Path) -> None:
119 """Write evaluation metrics as a one-row CSV file.
121 Args:
122 metrics: Metric names mapped to numeric scores.
123 output_csv: Destination CSV path.
124 """
125 output = Path(output_csv)
126 output.parent.mkdir(parents=True, exist_ok=True)
127 pd.DataFrame([metrics]).to_csv(output, index=False)