#!/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()