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

1"""Evaluation metrics for image clustering results.""" 

2 

3from pathlib import Path 

4 

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 

12 

13from .io import ImclusterIO 

14 

15console = Console() 

16 

17 

18def clustering_accuracy(expected: ArrayLike, predicted: ArrayLike) -> float: 

19 """Return clustering accuracy under the optimal class-label assignment. 

20 

21 Args: 

22 expected: Ground-truth class labels. 

23 predicted: Predicted cluster labels. 

24 

25 Returns: 

26 The fraction of samples assigned correctly after optimally matching 

27 cluster IDs to expected classes. 

28 

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") 

36 

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)) 

43 

44 

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. 

51 

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. 

56 

57 Returns: 

58 NMI, ARI, and optimally matched clustering accuracy. 

59 

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}") 

84 

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 ) 

92 

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 } 

100 

101 

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) 

116 

117 

118def write_evaluation(metrics: dict[str, float], output_csv: str | Path) -> None: 

119 """Write evaluation metrics as a one-row CSV file. 

120 

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)