chq1155 commited on
Commit
d1b1033
·
verified ·
1 Parent(s): 643b8e1

Remove old eval/aggregate_universal_v2.py (renamed to main_benchmark)

Browse files
Files changed (1) hide show
  1. eval/aggregate_universal_v2.py +0 -188
eval/aggregate_universal_v2.py DELETED
@@ -1,188 +0,0 @@
1
- #!/usr/bin/env python3
2
- """Aggregate per-(case, method) evals.tsv into universal_v2 main_table and pattern_summary."""
3
- from __future__ import annotations
4
- import csv
5
- import math
6
- import os
7
- from pathlib import Path
8
-
9
- import yaml
10
-
11
- # Bundled benchmark data (cases.yaml) ships in bench/universal_v2 alongside this
12
- # script; the per-(case,method) results tree (evals.tsv, *_summary.tsv) is supplied
13
- # by the user via SF_UV2_RESULTS (predictions are not distributed with the package).
14
- DATA_ROOT = Path(os.environ.get(
15
- "SF_UV2_DATA", Path(__file__).resolve().parents[1] / "bench" / "universal_v2"))
16
- RES_ROOT = Path(os.environ.get(
17
- "SF_UV2_RESULTS", Path(__file__).resolve().parents[1] / "bench" / "universal_v2" / "results"))
18
-
19
- METHODS = ["mosaic_raw", "gradient_raw", "af_cluster",
20
- "depth_matched_random", "diversity_matched_random", "fi_shuffled_control"]
21
-
22
-
23
- def load_cases() -> dict[str, dict]:
24
- cs = yaml.safe_load((DATA_ROOT / "cases.yaml").read_text())["cases"]
25
- return {c["case_id"]: c for c in cs}
26
-
27
-
28
- def is_id(case_id: str) -> bool:
29
- return case_id.startswith("SFB_ID_")
30
-
31
-
32
- def parse_tsv(path: Path) -> list[dict]:
33
- if not path.exists():
34
- return []
35
- with path.open() as f:
36
- return list(csv.DictReader(f, delimiter="\t"))
37
-
38
-
39
- def _f(x):
40
- try:
41
- v = float(x)
42
- if math.isnan(v):
43
- return None
44
- return v
45
- except Exception:
46
- return None
47
-
48
-
49
- def _hit(x):
50
- return x in ("1", "True", "true")
51
-
52
-
53
- def main():
54
- cases = load_cases()
55
- pattern_map = {"FS": "fold_switch_metamorphic",
56
- "AL": "allosteric_ligand_induced",
57
- "ID": "idp_idr_disorder_to_order",
58
- "OL": "oligomer_domain_swap"}
59
-
60
- main_rows = []
61
- for cid, case in cases.items():
62
- pattern = case["pattern"]
63
- pat_short = cid.split("_")[1] # FS / AL / ID / OL
64
- for method in METHODS:
65
- screen_summary = RES_ROOT / cid / method / "screen_summary.tsv"
66
- refine_summary = RES_ROOT / cid / method / "refine_summary.tsv"
67
- evals_tsv = RES_ROOT / cid / method / "refine_per_state" / "evals.tsv"
68
-
69
- screen_rows = parse_tsv(screen_summary)
70
- refine_rows = parse_tsv(refine_summary)
71
- eval_rows = parse_tsv(evals_tsv)
72
-
73
- n_screen = len(screen_rows)
74
- n_refine = len(refine_rows)
75
- n_eval = len(eval_rows)
76
-
77
- hit_a = sum(1 for r in eval_rows if _hit(r.get("state_a__hit_primary", "0")))
78
- best_rmsd_a_vals = [_f(r.get("state_a__rmsd_common_core_A"))
79
- for r in eval_rows]
80
- best_rmsd_a_vals = [v for v in best_rmsd_a_vals if v is not None]
81
- best_rmsd_a = min(best_rmsd_a_vals) if best_rmsd_a_vals else None
82
-
83
- if is_id(cid):
84
- hit_b = None
85
- best_rmsd_b = None
86
- else:
87
- hit_b = sum(1 for r in eval_rows if _hit(r.get("state_b__hit_primary", "0")))
88
- best_rmsd_b_vals = [_f(r.get("state_b__rmsd_common_core_A"))
89
- for r in eval_rows]
90
- best_rmsd_b_vals = [v for v in best_rmsd_b_vals if v is not None]
91
- best_rmsd_b = min(best_rmsd_b_vals) if best_rmsd_b_vals else None
92
-
93
- main_rows.append({
94
- "case_id": cid,
95
- "pattern": pattern,
96
- "pattern_short": pat_short,
97
- "method": method,
98
- "n_screen": n_screen,
99
- "n_refine": n_refine,
100
- "n_eval": n_eval,
101
- "hit_count_stateA": hit_a,
102
- "hit_count_stateB": "NA" if hit_b is None else hit_b,
103
- "hit_rate_stateA": f"{hit_a / n_eval:.4f}" if n_eval else "NA",
104
- "hit_rate_stateB": ("NA" if hit_b is None else
105
- (f"{hit_b / n_eval:.4f}" if n_eval else "NA")),
106
- "best_rmsd_stateA": f"{best_rmsd_a:.3f}" if best_rmsd_a is not None else "NA",
107
- "best_rmsd_stateB": ("NA" if best_rmsd_b is None else
108
- (f"{best_rmsd_b:.3f}" if best_rmsd_b is not None else "NA")),
109
- "total_inferences": n_screen + n_refine,
110
- })
111
-
112
- # main_table.csv
113
- out_main = RES_ROOT / "main_table.csv"
114
- out_main.parent.mkdir(parents=True, exist_ok=True)
115
- cols = ["case_id", "pattern", "pattern_short", "method",
116
- "n_screen", "n_refine", "n_eval",
117
- "hit_count_stateA", "hit_count_stateB",
118
- "hit_rate_stateA", "hit_rate_stateB",
119
- "best_rmsd_stateA", "best_rmsd_stateB",
120
- "total_inferences"]
121
- with out_main.open("w") as f:
122
- f.write(",".join(cols) + "\n")
123
- for r in main_rows:
124
- f.write(",".join(str(r[c]) for c in cols) + "\n")
125
- print(f"wrote {len(main_rows)} rows → {out_main}")
126
-
127
- # pattern_summary.csv: per (pattern, method) → mean hit_rate (A) and (B), n_cases
128
- pat_summary: dict[tuple[str, str], dict] = {}
129
- for r in main_rows:
130
- key = (r["pattern_short"], r["method"])
131
- s = pat_summary.setdefault(key, {"hit_rates_A": [], "hit_rates_B": [],
132
- "cases_with_data": 0})
133
- if r["n_eval"] and r["n_eval"] != 0 and r["hit_rate_stateA"] != "NA":
134
- try:
135
- s["hit_rates_A"].append(float(r["hit_rate_stateA"]))
136
- s["cases_with_data"] += 1
137
- except Exception:
138
- pass
139
- if r["hit_rate_stateB"] not in ("NA", ""):
140
- try:
141
- s["hit_rates_B"].append(float(r["hit_rate_stateB"]))
142
- except Exception:
143
- pass
144
-
145
- import statistics
146
- pat_rows = []
147
- for (pat, method), s in pat_summary.items():
148
- a = s["hit_rates_A"]
149
- b = s["hit_rates_B"]
150
- pat_rows.append({
151
- "pattern": pat,
152
- "method": method,
153
- "n_cases": s["cases_with_data"],
154
- "mean_hit_rate_stateA": f"{statistics.mean(a):.4f}" if a else "NA",
155
- "std_hit_rate_stateA": f"{statistics.pstdev(a):.4f}" if len(a) > 1 else "NA",
156
- "mean_hit_rate_stateB": f"{statistics.mean(b):.4f}" if b else "NA",
157
- "std_hit_rate_stateB": f"{statistics.pstdev(b):.4f}" if len(b) > 1 else "NA",
158
- })
159
-
160
- out_pat = RES_ROOT / "pattern_summary.csv"
161
- pcols = ["pattern", "method", "n_cases",
162
- "mean_hit_rate_stateA", "std_hit_rate_stateA",
163
- "mean_hit_rate_stateB", "std_hit_rate_stateB"]
164
- pat_rows.sort(key=lambda x: (x["pattern"], x["method"]))
165
- with out_pat.open("w") as f:
166
- f.write(",".join(pcols) + "\n")
167
- for r in pat_rows:
168
- f.write(",".join(str(r[c]) for c in pcols) + "\n")
169
- print(f"wrote {len(pat_rows)} rows → {out_pat}")
170
-
171
- # Print headline table
172
- print("\n=== Per-pattern × method (mean hit_rate_stateA) ===")
173
- pats = ["FS", "AL", "ID", "OL"]
174
- print(f"{'pattern':<8}" + "".join(f"{m:<28}" for m in METHODS))
175
- for p in pats:
176
- line = f"{p:<8}"
177
- for m in METHODS:
178
- row = next((r for r in pat_rows
179
- if r["pattern"] == p and r["method"] == m), None)
180
- if row:
181
- line += f"{row['mean_hit_rate_stateA']:<8} (n={row['n_cases']:<2}) "
182
- else:
183
- line += f"{'--':<28}"
184
- print(line)
185
-
186
-
187
- if __name__ == "__main__":
188
- main()