-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathpredict_cli.py
More file actions
110 lines (89 loc) · 4.33 KB
/
Copy pathpredict_cli.py
File metadata and controls
110 lines (89 loc) · 4.33 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
"""Unified CLI for running the shipped classifier on a test set.
Examples
--------
Default (v4 multi-scale champion, 0.6887 F1):
.venv/bin/python predict_cli.py --input /path/to/TEST_SET --output submission.csv
Prior v2 TTA champion (0.6562 F1, faster):
.venv/bin/python predict_cli.py --model models/ensemble_v2_tta --input /path/to/TEST_SET
Non-TTA fallback (faster ~8x):
.venv/bin/python predict_cli.py --model models/ensemble_v1 --input /path/to/TEST_SET
Single-encoder fallback:
.venv/bin/python predict_cli.py --model models/dinov2b_tiled_v1 --input /path/to/TEST_SET
"""
from __future__ import annotations
import argparse
import sys
import time
from pathlib import Path
import pandas as pd
def _load_predictor(model_dir: Path):
"""Load either TTAPredictor (for ensemble_v1_tta) or TearClassifier (for others)."""
meta_path = model_dir / "meta.json"
if not meta_path.exists():
sys.exit(f"ERROR: {meta_path} not found — is this a valid model bundle?")
# TTA bundles ship their own predict.py with TTAPredictor
tta_predict = model_dir / "predict.py"
if tta_predict.exists():
import importlib.util
spec = importlib.util.spec_from_file_location("tta_predict", tta_predict)
mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(mod)
if hasattr(mod, "TTAPredictor"):
print(f"[load] TTAPredictor from {model_dir}")
return mod.TTAPredictor.load(model_dir), "tta"
# Otherwise use standard TearClassifier
from teardrop.infer import TearClassifier
print(f"[load] TearClassifier from {model_dir}")
return TearClassifier.load(model_dir), "standard"
def main():
ap = argparse.ArgumentParser(description="Predict tear AFM classes on a directory.")
ap.add_argument("--model", default="models/ensemble_v4_multiscale",
help="Path to model bundle (default: v4 multi-scale champion, 0.6887 F1)")
ap.add_argument("--input", required=True, help="Directory with raw SPM scans (recursive)")
ap.add_argument("--output", default="submission.csv",
help="Output CSV path (default: submission.csv)")
ap.add_argument("--progress-every", type=int, default=10,
help="Print progress every N scans (default: 10)")
ap.add_argument("--input-format", default="spm", choices=["spm", "bmp"],
help="Input format: 'spm' = raw Bruker .NNN (default, full "
"accuracy); 'bmp' = 704x575 BMP previews (fallback, "
"degraded accuracy — see teardrop/bmp_infer.py)")
args = ap.parse_args()
model_dir = Path(args.model).resolve()
input_dir = Path(args.input).resolve()
output_path = Path(args.output).resolve()
if not model_dir.exists():
sys.exit(f"ERROR: model dir not found: {model_dir}")
if not input_dir.exists():
sys.exit(f"ERROR: input dir not found: {input_dir}")
if args.input_format == "bmp":
# Only v4 multi-scale has BMP fallback implemented; other bundles
# should already have been migrated to SPM by the organizer.
if "ensemble_v4_multiscale" not in str(model_dir):
print(f"[warn] --input-format bmp is only validated against "
f"models/ensemble_v4_multiscale; got {model_dir.name}")
from teardrop.bmp_infer import BmpPredictorV4
print(f"[load] BmpPredictorV4 from {model_dir} (BMP fallback path)")
clf = BmpPredictorV4.load(model_dir)
else:
clf, _ = _load_predictor(model_dir)
t0 = time.time()
df = clf.predict_directory(input_dir)
elapsed = time.time() - t0
# Ensure column order: file, predicted_class, prob_*
df = df.sort_values("file").reset_index(drop=True)
df.to_csv(output_path, index=False)
print(f"\n[done] {len(df)} scans predicted in {elapsed:.1f} s "
f"({elapsed / max(1, len(df)):.2f} s/scan)")
print(f"[saved] {output_path}")
# Summary
if "predicted_class" in df.columns:
print("\nPrediction distribution:")
for cls, n in df["predicted_class"].value_counts().items():
print(f" {cls:20s} {n:5d}")
if "error" in df.columns:
n_err = df["error"].notna().sum()
if n_err > 0:
print(f"\n⚠ {n_err} scans failed (see 'error' column in output CSV)")
if __name__ == "__main__":
main()