-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathclassification_benchmark.py
More file actions
1735 lines (1494 loc) · 78.2 KB
/
Copy pathclassification_benchmark.py
File metadata and controls
1735 lines (1494 loc) · 78.2 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
#!/usr/bin/env python3
"""
classification_benchmark.py
============================
k-NN and linear probe classification benchmark for DeepvEM embeddings.
This script does:
1. Loads embeddings.npz (from extract_embeddings.py)
2. Loads the dataset annotation table (metadata/dataset_metadata.csv)
to get coarse & fine labels
3. Maps each image → coarse_label + fine_label via the dataset name
4. Splits into train/test (stratified)
5. Runs k-NN classification (no training needed)
6. Runs linear probe (LogisticRegression on frozen features)
7. Optionally runs data-efficiency curves over reader-chosen training fractions
8. Saves results as JSON + prints paper-ready tables
Can benchmark multiple embeddings files (pretrained vs random init) side by side.
No hyperparameter that affects a reported number has a built-in default (see
--help): every one is required on the command line, so a result can always be
traced back to the exact settings that produced it.
Usage:
# Single model evaluation:
python classification_benchmark.py \
--embeddings /path/to/embeddings.npz \
--metadata metadata/dataset_metadata.csv \
--model_name "<model name>" \
--knn_k <k> --knn_k_values <k1> <k2> <k3> \
--test_fraction <fraction> --min_samples <n> --seed <seed> \
--linear_max_iter <max_iter> --linear_C <C>
# Compare pretrained vs random init:
python classification_benchmark.py \
--embeddings /path/to/pretrained_embeddings.npz /path/to/random_embeddings.npz \
--model_names "<model A>" "<model B>" \
--metadata metadata/dataset_metadata.csv \
--knn_k <k> --knn_k_values <k1> <k2> <k3> \
--test_fraction <fraction> --min_samples <n> --seed <seed> \
--linear_max_iter <max_iter> --linear_C <C>
# With data-efficiency curve (adds --de_fractions and --de_repeats):
python classification_benchmark.py \
--embeddings /path/to/embeddings.npz \
--metadata metadata/dataset_metadata.csv \
--knn_k <k> --knn_k_values <k1> <k2> <k3> \
--test_fraction <fraction> --min_samples <n> --seed <seed> \
--linear_max_iter <max_iter> --linear_C <C> \
--data_efficiency --de_fractions <f1> <f2> ... --de_repeats <n>
"""
import argparse
import csv
import json
import os
import re
import sys
import warnings
from collections import defaultdict
from pathlib import Path
import numpy as np
import pandas as pd
from sklearn.linear_model import LogisticRegression
from sklearn.neighbors import KNeighborsClassifier
from sklearn.model_selection import StratifiedShuffleSplit, GroupShuffleSplit
from sklearn.preprocessing import LabelEncoder, StandardScaler
from sklearn.metrics import (
accuracy_score, f1_score, classification_report,
confusion_matrix, top_k_accuracy_score
)
warnings.filterwarnings('ignore', category=UserWarning)
# Label granularities this benchmark knows how to evaluate. Used both to
# validate --label_levels and to recognise a label level when merging saved
# result JSONs (see _load_and_merge_jsons) -- a single source of truth so the
# two can't silently drift apart.
VALID_LABEL_LEVELS = ('coarse', 'fine', 'organism', 'modality')
# No hyperparameter here has a built-in numeric default (see main()'s argparse
# section): every value that would change a reported number is required from
# the command line, so a run can never silently pick up a number nobody chose.
# ============================================================================
# Embedding preprocessing
# ============================================================================
def preprocess_embeddings(X: np.ndarray, normalize: bool = True) -> np.ndarray:
"""
Preprocess embeddings before classification.
L2 normalization is CRITICAL for fair comparison between pretrained
and random init models. Why?
Random init encoders produce features with dataset-specific magnitude
patterns — different microscopes/staining produce different overall
intensity statistics, which translate to different feature vector norms.
A classifier can exploit these magnitude differences as a "shortcut"
to achieve high accuracy without learning any real biological structure.
L2 normalization removes this shortcut by projecting all embeddings
onto the unit hypersphere, forcing the classifier to use DIRECTIONAL
structure (angular relationships between features) rather than magnitude.
This is the standard practice in contrastive learning evaluation
(SimCLR, MoCo, DINO, EM-DINO all use L2-normalized features for
classification benchmarks).
"""
if normalize:
norms = np.linalg.norm(X, axis=1, keepdims=True)
norms = np.maximum(norms, 1e-8) # avoid division by zero
X = X / norms
return X
DEFAULT_METADATA_CSV = Path(__file__).resolve().parent / "metadata" / "dataset_metadata.csv"
# Columns are matched case-insensitively, accepting either the names used by
# metadata/dataset_metadata.csv or the longer ones from the original spreadsheet
# export, so both work without editing the table.
_COLUMN_ALIASES = {
'dataset_id': ('dataset_id',),
'organism': ('organism',),
'tissue_fine': ('tissue_fine', 'tissue_organ_fine'),
'tissue_coarse': ('tissue_coarse', 'tissue_organ_coarse'),
'fine_label': ('umap_fine', 'umap_label_fine'),
'coarse_label': ('umap_coarse', 'umap_label_coarse'),
'modality': ('modality',),
}
def load_metadata(path: str) -> list:
"""Read the dataset annotation table into a list of lower-cased dict rows.
Accepts .csv/.tsv (stdlib) and .xlsx/.xls (via pandas).
"""
path = Path(path)
if not path.is_file():
raise SystemExit(
f"[ERROR] Dataset metadata not found: {path}\n"
f" Pass --metadata /path/to/table.csv "
f"(schema: {DEFAULT_METADATA_CSV})"
)
if path.suffix.lower() in ('.xlsx', '.xls'):
try:
df = pd.read_excel(path, sheet_name='Combined_All_Datasets')
except Exception:
df = pd.read_excel(path, sheet_name=0)
rows = df.to_dict('records')
else:
delim = '\t' if path.suffix.lower() in ('.tsv', '.tab') else ','
with path.open(newline='', encoding='utf-8-sig') as f:
rows = list(csv.DictReader(f, delimiter=delim))
return [{(k or '').strip().lower(): ('' if v is None else str(v).strip())
for k, v in row.items()} for row in rows]
def build_dataset_label_map(metadata_rows: list) -> dict:
"""
Build mapping: dataset_id -> {fine_label, coarse_label, organism, ...}
The dataset_id must correspond to the per-image label stored in
embeddings.npz; resolve_label() below also does fuzzy/prefix matching.
"""
def pick(row, field):
for name in _COLUMN_ALIASES[field]:
val = row.get(name, '')
if val and val.lower() != 'nan':
return val
return ''
label_map = {}
for row in metadata_rows:
did = pick(row, 'dataset_id')
if not did:
continue
organism = pick(row, 'organism') or 'Unknown'
tissue_fine = pick(row, 'tissue_fine') or 'Unknown'
tissue_coarse = pick(row, 'tissue_coarse') or 'Unknown'
# The explicit umap_* columns win. When absent, derive "<organism> - <tissue>"
# so that the class is organism-aware -- otherwise e.g. mouse brain and
# Drosophila brain would collapse into a single "Brain" class.
def derive(explicit, tissue):
if explicit:
return explicit
if organism != 'Unknown' and tissue != 'Unknown':
return f'{organism} - {tissue}'
return tissue
label_map[did] = {
'fine_label': derive(pick(row, 'fine_label'), tissue_fine),
'coarse_label': derive(pick(row, 'coarse_label'), tissue_coarse),
'organism': organism,
'tissue_fine': tissue_fine,
'tissue_coarse': tissue_coarse,
'modality': pick(row, 'modality') or 'Unknown',
# Alternative names that should resolve to this dataset.
'aliases': [a.strip() for a in row.get('aliases', '').split(';') if a.strip()],
}
if not label_map:
raise SystemExit("[ERROR] No usable rows in the metadata table "
"(is the 'dataset_id' column present?)")
return label_map
def build_fuzzy_lookup(label_map: dict) -> dict:
"""
Build a fuzzy lookup: normalized_key -> dataset_id.
Normalization removes hyphens, underscores and spaces, and lowercases.
Alternative names from the table's `aliases` column are included, so a label
that bears no resemblance to the canonical id still resolves.
"""
def norm(s):
return s.replace('-', '').replace('_', '').replace(' ', '').lower()
fuzzy = {}
for did, entry in label_map.items():
fuzzy[norm(did)] = did
# Aliases are added second, and never overwrite a real dataset id.
for did, entry in label_map.items():
for alias in entry.get('aliases', []):
fuzzy.setdefault(norm(alias), did)
return fuzzy
def resolve_label(raw_label: str, label_map: dict, fuzzy_lookup: dict) -> dict:
"""
Resolve a raw label (from embeddings.npz) to its metadata entry.
Returns the metadata dict or None if not found.
"""
# Direct match
if raw_label in label_map:
return label_map[raw_label]
# Fuzzy match
normalized = raw_label.replace('-', '').replace('_', '').replace(' ', '').lower()
if normalized in fuzzy_lookup:
return label_map[fuzzy_lookup[normalized]]
# Partial match: try if raw_label is a prefix of any Dataset_ID
for did in label_map:
did_norm = did.replace('-', '').replace('_', '').replace(' ', '').lower()
if did_norm.startswith(normalized) or normalized.startswith(did_norm):
return label_map[did]
return None
# ============================================================================
# Core classification experiments
# ============================================================================
def run_knn(X_train, y_train, X_test, y_test, k_values):
"""
Run k-NN classification for multiple k values.
This is the simplest evaluation: no training at all.
Just finds the k nearest neighbors in the training set and votes.
Reports TOP-1 accuracy (like Cryo-IEF's "TOP1 k-NN" metric).
"""
results = {}
for k in k_values:
if k > len(X_train):
continue
knn = KNeighborsClassifier(n_neighbors=k, metric='cosine', n_jobs=-1)
knn.fit(X_train, y_train)
y_pred = knn.predict(X_test)
acc = accuracy_score(y_test, y_pred)
f1_macro = f1_score(y_test, y_pred, average='macro', zero_division=0)
f1_weighted = f1_score(y_test, y_pred, average='weighted', zero_division=0)
results[f'knn_k{k}'] = {
'accuracy': float(acc),
'f1_macro': float(f1_macro),
'f1_weighted': float(f1_weighted),
}
return results
def run_linear_probe(X_train, y_train, X_test, y_test,
max_iter, C, return_model=False):
"""
Run linear probe classification (LogisticRegression).
=== What is a linear probe? ===
A linear probe is the simplest possible classifier: a single linear layer
(no hidden layers, no nonlinearities). It takes the frozen encoder features
and learns a weight matrix W and bias b such that:
predicted_class = argmax(W @ feature_vector + b)
That's it. It's equivalent to drawing straight-line decision boundaries
in the feature space.
WHY does this matter? If a linear probe gets HIGH accuracy, it means the
pretrained features are already well-separated in feature space — the
different classes form distinct clusters that can be separated by hyperplanes.
This is a direct measure of FEATURE QUALITY from pretraining.
If a random init encoder gets low linear probe accuracy but pretrained gets
high accuracy, that proves pretraining learned meaningful representations.
WHY does it get high accuracy in papers? Because good pretraining
(contrastive learning, MAE, etc.) explicitly encourages the encoder to
produce features where similar things are close and different things are
far apart. A linear classifier then trivially separates them.
It's "easy" BY DESIGN — that's the point of the evaluation.
===
"""
# Normalize features (important for logistic regression)
scaler = StandardScaler()
X_train_scaled = scaler.fit_transform(X_train)
X_test_scaled = scaler.transform(X_test)
# L-BFGS solver handles multi-class well, with L2 regularization
clf = LogisticRegression(
max_iter=max_iter,
solver='lbfgs',
C=C, # Regularization strength (inverse)
n_jobs=-1,
)
clf.fit(X_train_scaled, y_train)
y_pred = clf.predict(X_test_scaled)
acc = accuracy_score(y_test, y_pred)
f1_macro = f1_score(y_test, y_pred, average='macro', zero_division=0)
f1_weighted = f1_score(y_test, y_pred, average='weighted', zero_division=0)
metrics = {
'accuracy': float(acc),
'f1_macro': float(f1_macro),
'f1_weighted': float(f1_weighted),
}
# Callers that also want the per-sample predictions (e.g. to build a
# classification report) take them from here rather than refitting.
return (metrics, y_pred) if return_model else metrics
def run_data_efficiency(X_train, y_train, X_test, y_test,
fractions, n_repeats, seed, knn_k,
linear_max_iter, linear_C, extra_knn_ks=None):
"""
Run BOTH k-NN and linear probe with varying amounts of training data.
This produces the data-efficiency curve.
At each fraction, we subsample the training set (stratified),
run both classifiers, and report accuracy. Repeated n_repeats times
for error bars (at 100% there's no variance since no subsampling).
k-NN data efficiency: fewer training points = fewer neighbors to look up.
This tests whether the feature space is structured enough that even a
handful of labeled examples suffice to classify new images.
extra_knn_ks: additional k values to also evaluate at each fraction,
stored as knn_k{k} keys alongside the primary 'knn' key.
"""
extra_knn_ks = [k for k in (extra_knn_ks or []) if k != knn_k]
results = {}
rng = np.random.RandomState(seed)
for frac in fractions:
lp_accs, lp_f1s = [], []
knn_accs, knn_f1s = [], []
extra_accs = {k: [] for k in extra_knn_ks}
extra_f1s = {k: [] for k in extra_knn_ks}
# At 100% there is no subsampling, so every repeat is identical: run it
# once rather than n_repeats times (the std is 0 by construction).
reps = 1 if frac >= 1.0 else n_repeats
n_sub_actual = len(X_train)
for rep in range(reps):
if frac >= 1.0:
X_sub, y_sub = X_train, y_train
else:
n_sub = max(1, int(len(X_train) * frac))
n_sub_actual = n_sub
try:
sss = StratifiedShuffleSplit(n_splits=1, train_size=n_sub,
random_state=seed + rep)
idx, _ = next(sss.split(X_train, y_train))
X_sub = X_train[idx]
y_sub = y_train[idx]
except ValueError:
idx = rng.choice(len(X_train), n_sub, replace=False)
X_sub = X_train[idx]
y_sub = y_train[idx]
if len(np.unique(y_sub)) < 2:
continue
# Linear probe
lp_res = run_linear_probe(X_sub, y_sub, X_test, y_test,
max_iter=linear_max_iter, C=linear_C)
lp_accs.append(lp_res['accuracy'])
lp_f1s.append(lp_res['f1_macro'])
# Primary k-NN (adjust k if fewer samples than k)
effective_k = min(knn_k, len(X_sub))
if effective_k < 1:
effective_k = 1
knn = KNeighborsClassifier(n_neighbors=effective_k, metric='cosine', n_jobs=-1)
knn.fit(X_sub, y_sub)
y_pred_knn = knn.predict(X_test)
knn_accs.append(float(accuracy_score(y_test, y_pred_knn)))
knn_f1s.append(float(f1_score(y_test, y_pred_knn, average='macro', zero_division=0)))
# Extra k values
for ek in extra_knn_ks:
eff_ek = min(ek, len(X_sub))
if eff_ek < 1:
eff_ek = 1
knn_ek = KNeighborsClassifier(n_neighbors=eff_ek, metric='cosine', n_jobs=-1)
knn_ek.fit(X_sub, y_sub)
y_pred_ek = knn_ek.predict(X_test)
extra_accs[ek].append(float(accuracy_score(y_test, y_pred_ek)))
extra_f1s[ek].append(float(f1_score(y_test, y_pred_ek, average='macro', zero_division=0)))
if lp_accs:
frac_pct = f"{int(frac * 100)}%"
primary_knn_entry = {
'accuracy_mean': float(np.mean(knn_accs)),
'accuracy_std': float(np.std(knn_accs)),
'f1_macro_mean': float(np.mean(knn_f1s)),
'f1_macro_std': float(np.std(knn_f1s)),
}
results[frac_pct] = {
'linear_probe': {
'accuracy_mean': float(np.mean(lp_accs)),
'accuracy_std': float(np.std(lp_accs)),
'f1_macro_mean': float(np.mean(lp_f1s)),
'f1_macro_std': float(np.std(lp_f1s)),
},
'knn': primary_knn_entry,
f'knn_k{knn_k}': primary_knn_entry, # named key for explicit lookup
'n_train_samples': int(n_sub_actual),
}
for ek in extra_knn_ks:
if extra_accs[ek]:
results[frac_pct][f'knn_k{ek}'] = {
'accuracy_mean': float(np.mean(extra_accs[ek])),
'accuracy_std': float(np.std(extra_accs[ek])),
'f1_macro_mean': float(np.mean(extra_f1s[ek])),
'f1_macro_std': float(np.std(extra_f1s[ek])),
}
return results
# ============================================================================
# Main pipeline
# ============================================================================
def evaluate_embeddings(
embeddings: np.ndarray,
labels_raw: np.ndarray,
names: np.ndarray,
label_map: dict,
fuzzy_lookup: dict,
knn_k: int,
knn_k_values: list,
test_fraction: float,
min_samples_per_class: int,
seed: int,
linear_max_iter: int,
linear_C: float,
label_level: str = 'coarse',
data_efficiency: bool = False,
normalize: bool = True,
extra_knn_ks: list = None,
de_fractions: list = None,
de_repeats: int = None,
group_split: bool = False,
):
"""
Full evaluation pipeline for one set of embeddings.
Returns a results dict with k-NN, linear probe, and optionally data-efficiency.
"""
# Step 0: L2-normalize embeddings
embeddings = preprocess_embeddings(embeddings, normalize=normalize)
# Step 1: Map raw labels to coarse/fine classification labels
mapped_labels = []
unmapped_count = 0
unmapped_datasets = set()
for raw in labels_raw:
raw_str = str(raw)
meta = resolve_label(raw_str, label_map, fuzzy_lookup)
if meta is not None:
# Map label_level to the appropriate metadata field
level_to_key = {
'coarse': 'coarse_label',
'fine': 'fine_label',
'organism': 'organism',
'modality': 'modality',
}
key = level_to_key.get(label_level, 'coarse_label')
mapped_labels.append(meta[key])
else:
mapped_labels.append(None)
unmapped_count += 1
unmapped_datasets.add(raw_str)
mapped_labels = np.array(mapped_labels)
# Report mapping. mapped_labels is an object array, so compare elementwise
# rather than with `!=`, which broadcasts oddly against None.
valid_mask = np.array([l is not None for l in mapped_labels])
print(f"\n Label mapping ({label_level}):")
print(f" Mapped: {valid_mask.sum()} / {len(mapped_labels)}")
if unmapped_count > 0:
print(f" Unmapped ({unmapped_count}): {sorted(unmapped_datasets)[:10]}")
# Filter to valid only. The source-dataset id of every kept sample is carried
# along so that it can be used as a grouping key for the split below.
labels_raw_arr = np.asarray([str(r) for r in labels_raw])
embeddings_valid = embeddings[valid_mask]
labels_valid = mapped_labels[valid_mask].astype(str)
groups_valid = labels_raw_arr[valid_mask]
# Step 2: Filter classes with too few samples
unique, counts = np.unique(labels_valid, return_counts=True)
keep_classes = unique[counts >= min_samples_per_class]
class_mask = np.isin(labels_valid, keep_classes)
embeddings_valid = embeddings_valid[class_mask]
labels_valid = labels_valid[class_mask]
groups_valid = groups_valid[class_mask]
print(f" Classes with >= {min_samples_per_class} samples: {len(keep_classes)} / {len(unique)}")
print(f" Total samples after filtering: {len(labels_valid)}")
# Encode labels as integers
le = LabelEncoder()
y = le.fit_transform(labels_valid)
X = embeddings_valid
num_classes = len(le.classes_)
print(f" Number of classes: {num_classes}")
# Print class distribution
unique_y, counts_y = np.unique(y, return_counts=True)
print(f"\n Class distribution:")
for cls_idx, cnt in sorted(zip(unique_y, counts_y), key=lambda x: -x[1]):
print(f" [{cls_idx:2d}] {le.classes_[cls_idx]}: {cnt}")
# Step 3: Train/test split
#
# Default (--group_split off): a stratified split over individual volumes.
# Note that several volumes are typically cropped from the SAME source
# dataset, so train and test can contain highly correlated neighbours. This
# measures how well the features separate classes, but it is optimistic as a
# measure of generalization to an unseen dataset.
#
# --group_split on: whole source datasets are held out, so no dataset
# contributes to both train and test. This is the stricter protocol; it
# requires each class to be covered by at least two source datasets.
if group_split:
n_groups = len(np.unique(groups_valid))
gss = GroupShuffleSplit(n_splits=1, test_size=test_fraction, random_state=seed)
train_idx, test_idx = next(gss.split(X, y, groups=groups_valid))
held_out = sorted({str(g) for g in groups_valid[test_idx]})
print(f"\n Group-aware split: {len(held_out)} / {n_groups} source datasets held out")
print(f" Held out: {held_out[:8]}{' ...' if len(held_out) > 8 else ''}")
# Holding out whole datasets can strip every training example of a class.
# Such a class is unlearnable, and sklearn would otherwise fail later with
# an opaque "solver needs samples of at least 2 classes" error. Drop those
# test samples and report it, rather than crashing cryptically.
trainable = set(np.unique(y[train_idx]).tolist())
missing = sorted(set(np.unique(y).tolist()) - trainable)
if missing:
keep = np.isin(y[test_idx], list(trainable))
dropped = int((~keep).sum())
test_idx = test_idx[keep]
print(f" WARNING: {len(missing)} class(es) have no training data after the "
f"group split: {[le.classes_[c] for c in missing]}")
print(f" Dropped {dropped} test sample(s) of those classes. To avoid this, "
f"lower --test_fraction, use a coarser --label_levels, or ensure every "
f"class is covered by at least two source datasets.")
if len(trainable) < 2:
raise SystemExit(
f"[ERROR] --group_split left only {len(trainable)} class(es) in the training "
f"split; at least 2 are required.\n"
f" This corpus has too few source datasets per class for a dataset-level "
f"split at label level '{label_level}'.\n"
f" Use a coarser --label_levels, lower --test_fraction, or omit --group_split."
)
if len(test_idx) == 0:
raise SystemExit("[ERROR] --group_split left an empty test set. "
"Lower --test_fraction or omit --group_split.")
else:
sss = StratifiedShuffleSplit(n_splits=1, test_size=test_fraction, random_state=seed)
train_idx, test_idx = next(sss.split(X, y))
X_train, X_test = X[train_idx], X[test_idx]
y_train, y_test = y[train_idx], y[test_idx]
print(f"\n Train: {len(X_train)}, Test: {len(X_test)}")
# Step 4: Run benchmarks
results = {
'num_classes': num_classes,
'num_train': len(X_train),
'num_test': len(X_test),
'embedding_dim': X.shape[1],
'class_names': list(le.classes_),
'label_level': label_level,
'knn_k': knn_k,
'split': 'group' if group_split else 'stratified',
'test_source_datasets': sorted({str(g) for g in groups_valid[test_idx]}),
}
# k-NN. The primary k must be present in the evaluated set, otherwise the
# summary/plotting lookup (_get_knn_key) silently falls back to another k.
k_values = sorted(set(knn_k_values) | {knn_k})
print(f"\n Running k-NN classification (k={k_values}, primary k={knn_k})...")
knn_results = run_knn(X_train, y_train, X_test, y_test, k_values=k_values)
results['knn'] = knn_results
best_knn_k = max(knn_results, key=lambda k: knn_results[k]['accuracy'])
print(f" Best k-NN: {best_knn_k} -> acc={knn_results[best_knn_k]['accuracy']:.4f}")
# Linear probe. Keep the fitted predictions so the classification report
# below describes exactly this model instead of an independently refit one.
print(f" Running linear probe...")
lp_results, y_pred = run_linear_probe(X_train, y_train, X_test, y_test,
max_iter=linear_max_iter, C=linear_C,
return_model=True)
results['linear_probe'] = lp_results
print(f" Linear probe: acc={lp_results['accuracy']:.4f}, F1={lp_results['f1_macro']:.4f}")
# Data efficiency
if data_efficiency:
print(f" Running data-efficiency curve (k-NN k={knn_k} + linear probe)...")
de_results = run_data_efficiency(X_train, y_train, X_test, y_test,
fractions=de_fractions,
n_repeats=de_repeats,
seed=seed,
knn_k=knn_k,
extra_knn_ks=extra_knn_ks,
linear_max_iter=linear_max_iter,
linear_C=linear_C)
results['data_efficiency'] = de_results
for frac, vals in de_results.items():
lp_a = vals['linear_probe']['accuracy_mean']
lp_s = vals['linear_probe']['accuracy_std']
knn_a = vals['knn']['accuracy_mean']
knn_s = vals['knn']['accuracy_std']
print(f" {frac:>4s}: Linear={lp_a*100:5.1f}±{lp_s*100:.1f}% k-NN={knn_a*100:5.1f}±{knn_s*100:.1f}% (n={vals['n_train_samples']})")
# Detailed per-class report for the linear probe fitted above. (This used to
# refit an identical LogisticRegression, doubling the cost of the most
# expensive step and risking a report that disagreed with the reported
# metrics if the two fits were ever configured differently.)
report = classification_report(y_test, y_pred, target_names=le.classes_,
output_dict=True, zero_division=0)
results['classification_report'] = report
return results
def _get_knn_key(results: dict) -> str:
"""Get the primary knn key from results, respecting stored knn_k."""
knn_results = results.get('knn', {})
if 'knn_k' in results:
key = f"knn_k{results['knn_k']}"
if key in knn_results:
return key
# 'knn_k' missing or its k wasn't evaluated (older/malformed result file):
# fall back to the smallest k actually present, rather than guessing one.
if knn_results:
return min(knn_results, key=lambda k: int(k.replace('knn_k', '')))
# No k-NN results at all; any string is safe here since callers only ever
# do `.get(key, {})` against this key.
return 'knn_k'
def print_comparison_table(all_results: dict, label_level: str):
"""Print a paper-ready comparison table."""
print(f"\n{'='*90}")
print(f" CLASSIFICATION BENCHMARK RESULTS ({label_level.upper()} labels)")
print(f"{'='*90}")
# Determine k from first model
first_results = list(all_results.values())[0]
knn_key = _get_knn_key(first_results)
k_val = knn_key.replace('knn_k', '')
# Header
header = f"{'Model':<35} {'k-NN(k='+k_val+') Acc':<14} {'k-NN F1':<10} {'LP Acc':<10} {'LP F1':<10}"
print(header)
print("-" * 90)
for model_name, results in all_results.items():
knn_acc = results.get('knn', {}).get(_get_knn_key(results), {}).get('accuracy', -1)
knn_f1 = results.get('knn', {}).get(_get_knn_key(results), {}).get('f1_macro', -1)
lp_acc = results.get('linear_probe', {}).get('accuracy', -1)
lp_f1 = results.get('linear_probe', {}).get('f1_macro', -1)
knn_acc_s = f"{knn_acc*100:.2f}%" if knn_acc >= 0 else "N/A"
knn_f1_s = f"{knn_f1*100:.2f}%" if knn_f1 >= 0 else "N/A"
lp_acc_s = f"{lp_acc*100:.2f}%" if lp_acc >= 0 else "N/A"
lp_f1_s = f"{lp_f1*100:.2f}%" if lp_f1 >= 0 else "N/A"
print(f"{model_name:<35} {knn_acc_s:<10} {knn_f1_s:<10} {lp_acc_s:<10} {lp_f1_s:<10}")
print("-" * 90)
# Data efficiency table if available. The fractions actually present are
# read from the results rather than assumed, since --de_fractions is
# reader-chosen and may not match any particular fixed list.
has_de = any('data_efficiency' in r for r in all_results.values())
if has_de:
frac_keys = set()
for r in all_results.values():
frac_keys.update(r.get('data_efficiency', {}).keys())
fracs = sorted(frac_keys, key=lambda s: int(s.replace('%', '')))
for method, method_label in [('knn', 'k-NN'), ('linear_probe', 'Linear Probe')]:
for metric, metric_label in [('accuracy_mean', 'Accuracy'), ('f1_macro_mean', 'F1 Macro')]:
std_key = metric.replace('_mean', '_std')
print(f"\n Data Efficiency — {method_label} {metric_label} (%)")
print(f" {'Model':<30}", end="")
for f in fracs:
print(f" {f:>10}", end="")
print()
print(" " + "-" * 90)
for model_name, results in all_results.items():
de = results.get('data_efficiency', {})
print(f" {model_name:<30}", end="")
for f in fracs:
if f in de and method in de[f]:
val = de[f][method][metric]
std = de[f][method][std_key]
print(f" {val*100:>5.1f}±{std*100:.1f}", end="")
else:
print(f" {'N/A':>10}", end="")
print()
print()
# ============================================================================
# Figure plotting
# ============================================================================
def plot_results(all_output: dict, output_prefix: str = 'classification_figure',
show_std: bool = True):
"""
Generate figures from benchmark results.
Produces:
1. Data-efficiency line plots — one per label level, with k-NN and the
linear probe side by side
2. Summary bar chart of k-NN and linear probe at the full training set
Style: clean, no grid, muted but distinguishable colors, thin lines,
small markers.
"""
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from matplotlib.ticker import MaxNLocator
# ── Plot style ────────────────────────────────────────────────
plt.rcParams.update({
'font.family': 'sans-serif',
'font.sans-serif': ['Arial', 'Helvetica', 'DejaVu Sans'],
'font.size': 8,
'axes.titlesize': 9,
'axes.labelsize': 8,
'xtick.labelsize': 7,
'ytick.labelsize': 7,
'legend.fontsize': 7,
'figure.dpi': 300,
'savefig.dpi': 300,
'savefig.bbox': 'tight',
'axes.linewidth': 0.6,
'xtick.major.width': 0.6,
'ytick.major.width': 0.6,
'xtick.major.size': 3,
'ytick.major.size': 3,
'lines.linewidth': 1.2,
'lines.markersize': 4,
})
# Color palette — distinct and colorblind-friendly
COLORS = [
'#2166AC', # deep blue
'#B2182B', # deep red
'#1B7837', # dark green
'#FF7F00', # orange
'#984EA3', # purple
'#A65628', # brown
'#377EB8', # lighter blue
'#E41A1C', # bright red
]
MARKERS = ['o', 's', '^', 'D', 'v', 'P', 'X', 'h']
# Alternating thick/thin widths so a thick line is still visible beneath a thin one
LINEWIDTHS = [0.7, 0.7, 0.7, 0.7, 0.7, 0.7, 0.7, 0.7]
# Clearly distinct dash patterns using explicit on/off tuples
LINESTYLES = [
(0, ()), # solid
(0, (6, 2)), # long dash
(0, (3, 1, 1, 1)), # dash-dot
(0, (1, 1)), # dense dots
(0, (5, 1, 1, 1, 1, 1)), # dash-dot-dot
(0, (4, 2)), # medium dash
(0, (2, 1)), # short dash
(0, (1, 2)), # sparse dots
]
# Parse the fraction strings into numeric x-values
def frac_str_to_num(s):
return int(s.replace('%', ''))
for label_level, level_results in all_output.items():
model_names = list(level_results.keys())
# Check if data_efficiency exists
has_de = any('data_efficiency' in r for r in level_results.values())
if has_de:
# ── Figure 1: Data Efficiency Curves (2×2: metric × method) ──
# Top row: Accuracy, Bottom row: F1 Macro
# Left col: k-NN, Right col: Linear Probe
# Determine k from first model
first_r = level_results[model_names[0]]
plot_k = first_r.get('knn_k', 5)
knn_title = f'k-NN (k={plot_k})'
grid = [
# (row, col, method_key, metric_key, title, ylabel)
(0, 0, 'knn', 'accuracy_mean', 'accuracy_std', knn_title, 'Accuracy (%)'),
(0, 1, 'linear_probe', 'accuracy_mean', 'accuracy_std', 'Linear Probe', 'Accuracy (%)'),
(1, 0, 'knn', 'f1_macro_mean', 'f1_macro_std', knn_title, 'F1 Macro (%)'),
(1, 1, 'linear_probe', 'f1_macro_mean', 'f1_macro_std', 'Linear Probe', 'F1 Macro (%)'),
]
# Adaptive sizing: wider figure + external legend for many models
n_models = len(model_names)
fig_w = 6.0 if n_models <= 4 else 6.5
fig_h = 4.8 if n_models <= 4 else 5.2
fig, axes = plt.subplots(2, 2, figsize=(fig_w, fig_h))
# Adaptive marker size: smaller when many models
ms = max(3, 6 - n_models * 0.4)
lw = max(1.0, 1.6 - n_models * 0.08)
# Collect all y-values per panel to set smart y-limits
panel_y_data = {}
for row, col, method_key, metric_key, std_key, title, ylabel in grid:
ax = axes[row, col]
key = (row, col)
panel_y_data[key] = []
for i, model_name in enumerate(model_names):
de = level_results[model_name].get('data_efficiency', {})
if not de:
continue
fracs_str = sorted(de.keys(), key=frac_str_to_num)
x_vals = [frac_str_to_num(f) for f in fracs_str]
y_vals = [de[f][method_key][metric_key] * 100 for f in fracs_str]
y_errs = [de[f][method_key][std_key] * 100 for f in fracs_str]
panel_y_data[key].extend(y_vals)
color = COLORS[i % len(COLORS)]
marker = MARKERS[i % len(MARKERS)]
ls = LINESTYLES[i % len(LINESTYLES)]
lw_i = LINEWIDTHS[i % len(LINEWIDTHS)]
if show_std:
ax.errorbar(x_vals, y_vals, yerr=y_errs,
color=color, marker=marker, markersize=ms,
markerfacecolor=color, markeredgecolor='white',
markeredgewidth=0.5, linewidth=lw_i,
linestyle=ls,
capsize=2, capthick=0.5,
label=model_name if row == 0 else None,
zorder=3 + i * 0.1)
else:
ax.plot(x_vals, y_vals,
color=color, marker=marker, markersize=ms,
markerfacecolor=color, markeredgecolor='white',
markeredgewidth=0.5, linewidth=lw_i,
linestyle=ls,
label=model_name if row == 0 else None,
zorder=3 + i * 0.1)
# Smart y-axis: zoom to data range
if panel_y_data[key]:
y_min_data = min(panel_y_data[key])
y_max_data = max(panel_y_data[key])
y_range = y_max_data - y_min_data
padding = max(y_range * 0.15, 2.0)
ax.set_ylim(max(0, y_min_data - padding), min(100, y_max_data + padding))
ax.spines['top'].set_visible(False)
ax.spines['right'].set_visible(False)
ax.tick_params(direction='out')
ax.set_xlim(-3, 105)
ax.xaxis.set_major_locator(MaxNLocator(integer=True, nbins=6))
if row == 0:
ax.set_title(title, fontweight='semibold')
if row == 1:
ax.set_xlabel('Training Data (%)')
if col == 0:
ax.set_ylabel(ylabel)
# Legend — adaptive placement
handles, labels = axes[0, 0].get_legend_handles_labels()
n_legend = len(labels)
if n_legend <= 3:
axes[0, 1].legend(handles, labels, loc='lower right',
frameon=False, handletextpad=0.3, borderpad=0.2,
fontsize=7)
else:
# External legend below for many models — cleaner
ncol = min(n_legend, 3)
fig.legend(handles, labels, loc='lower center',
ncol=ncol, frameon=False,
bbox_to_anchor=(0.5, -0.06),
handletextpad=0.3, columnspacing=1.0,
fontsize=6.5)
n_classes = level_results[model_names[0]].get('num_classes', '?')
fig.suptitle(f'Classification — {label_level.capitalize()} Labels ({n_classes} classes)',
fontsize=9, fontweight='semibold', y=1.01)
plt.tight_layout()
fname = f'{output_prefix}_data_efficiency_{label_level}.png'
fig.savefig(fname, bbox_inches='tight', pad_inches=0.1)
plt.close(fig)
print(f" Saved: {fname}")
# ── Figure 2: Summary bar chart (Accuracy + F1) ─────────────
fig, axes_bar = plt.subplots(1, 2, figsize=(6.0, 2.5))
n_models = len(model_names)
x = np.arange(n_models)
bar_width = 0.35
for ax_idx, (metric_label, metric_key_knn, metric_key_lp) in enumerate([
('Accuracy (%)', 'accuracy', 'accuracy'),
('F1 Macro (%)', 'f1_macro', 'f1_macro'),
]):
ax = axes_bar[ax_idx]
knn_vals = []
lp_vals = []
for model_name in model_names:
r = level_results[model_name]
knn_vals.append(r.get('knn', {}).get(_get_knn_key(r), {}).get(metric_key_knn, 0) * 100)
lp_vals.append(r.get('linear_probe', {}).get(metric_key_lp, 0) * 100)
# Label with the k actually evaluated; a fixed "k=5" here mislabels
# every figure produced with a different --knn_k.
knn_k_used = level_results[model_names[0]].get('knn_k', '?') if model_names else '?'
bars1 = ax.bar(x - bar_width / 2, knn_vals, bar_width,
color=COLORS[0], edgecolor='white', linewidth=0.5,
label=f'k-NN (k={knn_k_used})' if ax_idx == 0 else None,
zorder=3, alpha=0.85)
bars2 = ax.bar(x + bar_width / 2, lp_vals, bar_width,
color=COLORS[1], edgecolor='white', linewidth=0.5,
label='Linear Probe' if ax_idx == 0 else None, zorder=3, alpha=0.85)
for bar in bars1:
h = bar.get_height()
ax.text(bar.get_x() + bar.get_width() / 2, h + 0.5,
f'{h:.1f}', ha='center', va='bottom', fontsize=5.5, color=COLORS[0])
for bar in bars2:
h = bar.get_height()
ax.text(bar.get_x() + bar.get_width() / 2, h + 0.5,
f'{h:.1f}', ha='center', va='bottom', fontsize=5.5, color=COLORS[1])
ax.set_xticks(x)
ax.set_xticklabels(model_names, rotation=20, ha='right', fontsize=6)
ax.set_ylabel(metric_label)
ax.spines['top'].set_visible(False)
ax.spines['right'].set_visible(False)
ax.tick_params(direction='out')
y_min = max(0, min(knn_vals + lp_vals) - 10)
ax.set_ylim(y_min, 105)
axes_bar[0].legend(frameon=False, fontsize=6.5, loc='lower right')