File size: 7,432 Bytes
bd082fe | 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 | #!/usr/bin/env python3
"""Aggregate per-(case, method) evals.tsv into main_benchmark main_table and pattern_summary."""
from __future__ import annotations
import csv
import math
import os
from pathlib import Path
import yaml
# Bundled benchmark data (cases.yaml) ships in bench/main_benchmark alongside this
# script; the per-(case,method) results tree (evals.tsv, *_summary.tsv) is supplied
# by the user via SF_MAIN_BENCH_RESULTS (predictions are not distributed with the package).
DATA_ROOT = Path(os.environ.get(
"SF_MAIN_BENCH_DATA", Path(__file__).resolve().parents[1] / "bench" / "main_benchmark"))
RES_ROOT = Path(os.environ.get(
"SF_MAIN_BENCH_RESULTS", Path(__file__).resolve().parents[1] / "bench" / "main_benchmark" / "results"))
METHODS = ["mosaic_raw", "gradient_raw", "af_cluster",
"depth_matched_random", "diversity_matched_random", "fi_shuffled_control"]
def load_cases() -> dict[str, dict]:
cs = yaml.safe_load((DATA_ROOT / "cases.yaml").read_text())["cases"]
return {c["case_id"]: c for c in cs}
def is_id(case_id: str) -> bool:
return case_id.startswith("SFB_ID_")
def parse_tsv(path: Path) -> list[dict]:
if not path.exists():
return []
with path.open() as f:
return list(csv.DictReader(f, delimiter="\t"))
def _f(x):
try:
v = float(x)
if math.isnan(v):
return None
return v
except Exception:
return None
def _hit(x):
return x in ("1", "True", "true")
def main():
cases = load_cases()
pattern_map = {"FS": "fold_switch_metamorphic",
"AL": "allosteric_ligand_induced",
"ID": "idp_idr_disorder_to_order",
"OL": "oligomer_domain_swap"}
main_rows = []
for cid, case in cases.items():
pattern = case["pattern"]
pat_short = cid.split("_")[1] # FS / AL / ID / OL
for method in METHODS:
screen_summary = RES_ROOT / cid / method / "screen_summary.tsv"
refine_summary = RES_ROOT / cid / method / "refine_summary.tsv"
evals_tsv = RES_ROOT / cid / method / "refine_per_state" / "evals.tsv"
screen_rows = parse_tsv(screen_summary)
refine_rows = parse_tsv(refine_summary)
eval_rows = parse_tsv(evals_tsv)
n_screen = len(screen_rows)
n_refine = len(refine_rows)
n_eval = len(eval_rows)
hit_a = sum(1 for r in eval_rows if _hit(r.get("state_a__hit_primary", "0")))
best_rmsd_a_vals = [_f(r.get("state_a__rmsd_common_core_A"))
for r in eval_rows]
best_rmsd_a_vals = [v for v in best_rmsd_a_vals if v is not None]
best_rmsd_a = min(best_rmsd_a_vals) if best_rmsd_a_vals else None
if is_id(cid):
hit_b = None
best_rmsd_b = None
else:
hit_b = sum(1 for r in eval_rows if _hit(r.get("state_b__hit_primary", "0")))
best_rmsd_b_vals = [_f(r.get("state_b__rmsd_common_core_A"))
for r in eval_rows]
best_rmsd_b_vals = [v for v in best_rmsd_b_vals if v is not None]
best_rmsd_b = min(best_rmsd_b_vals) if best_rmsd_b_vals else None
main_rows.append({
"case_id": cid,
"pattern": pattern,
"pattern_short": pat_short,
"method": method,
"n_screen": n_screen,
"n_refine": n_refine,
"n_eval": n_eval,
"hit_count_stateA": hit_a,
"hit_count_stateB": "NA" if hit_b is None else hit_b,
"hit_rate_stateA": f"{hit_a / n_eval:.4f}" if n_eval else "NA",
"hit_rate_stateB": ("NA" if hit_b is None else
(f"{hit_b / n_eval:.4f}" if n_eval else "NA")),
"best_rmsd_stateA": f"{best_rmsd_a:.3f}" if best_rmsd_a is not None else "NA",
"best_rmsd_stateB": ("NA" if best_rmsd_b is None else
(f"{best_rmsd_b:.3f}" if best_rmsd_b is not None else "NA")),
"total_inferences": n_screen + n_refine,
})
# main_table.csv
out_main = RES_ROOT / "main_table.csv"
out_main.parent.mkdir(parents=True, exist_ok=True)
cols = ["case_id", "pattern", "pattern_short", "method",
"n_screen", "n_refine", "n_eval",
"hit_count_stateA", "hit_count_stateB",
"hit_rate_stateA", "hit_rate_stateB",
"best_rmsd_stateA", "best_rmsd_stateB",
"total_inferences"]
with out_main.open("w") as f:
f.write(",".join(cols) + "\n")
for r in main_rows:
f.write(",".join(str(r[c]) for c in cols) + "\n")
print(f"wrote {len(main_rows)} rows → {out_main}")
# pattern_summary.csv: per (pattern, method) → mean hit_rate (A) and (B), n_cases
pat_summary: dict[tuple[str, str], dict] = {}
for r in main_rows:
key = (r["pattern_short"], r["method"])
s = pat_summary.setdefault(key, {"hit_rates_A": [], "hit_rates_B": [],
"cases_with_data": 0})
if r["n_eval"] and r["n_eval"] != 0 and r["hit_rate_stateA"] != "NA":
try:
s["hit_rates_A"].append(float(r["hit_rate_stateA"]))
s["cases_with_data"] += 1
except Exception:
pass
if r["hit_rate_stateB"] not in ("NA", ""):
try:
s["hit_rates_B"].append(float(r["hit_rate_stateB"]))
except Exception:
pass
import statistics
pat_rows = []
for (pat, method), s in pat_summary.items():
a = s["hit_rates_A"]
b = s["hit_rates_B"]
pat_rows.append({
"pattern": pat,
"method": method,
"n_cases": s["cases_with_data"],
"mean_hit_rate_stateA": f"{statistics.mean(a):.4f}" if a else "NA",
"std_hit_rate_stateA": f"{statistics.pstdev(a):.4f}" if len(a) > 1 else "NA",
"mean_hit_rate_stateB": f"{statistics.mean(b):.4f}" if b else "NA",
"std_hit_rate_stateB": f"{statistics.pstdev(b):.4f}" if len(b) > 1 else "NA",
})
out_pat = RES_ROOT / "pattern_summary.csv"
pcols = ["pattern", "method", "n_cases",
"mean_hit_rate_stateA", "std_hit_rate_stateA",
"mean_hit_rate_stateB", "std_hit_rate_stateB"]
pat_rows.sort(key=lambda x: (x["pattern"], x["method"]))
with out_pat.open("w") as f:
f.write(",".join(pcols) + "\n")
for r in pat_rows:
f.write(",".join(str(r[c]) for c in pcols) + "\n")
print(f"wrote {len(pat_rows)} rows → {out_pat}")
# Print headline table
print("\n=== Per-pattern × method (mean hit_rate_stateA) ===")
pats = ["FS", "AL", "ID", "OL"]
print(f"{'pattern':<8}" + "".join(f"{m:<28}" for m in METHODS))
for p in pats:
line = f"{p:<8}"
for m in METHODS:
row = next((r for r in pat_rows
if r["pattern"] == p and r["method"] == m), None)
if row:
line += f"{row['mean_hit_rate_stateA']:<8} (n={row['n_cases']:<2}) "
else:
line += f"{'--':<28}"
print(line)
if __name__ == "__main__":
main()
|