"""Batch-evaluate every `*_unrelaxed_rank_*.pdb` under a directory tree. For each PDB, runs evaluate_prediction.evaluate and emits one JSON sidecar next to the PDB (`_eval.json`) plus an aggregated TSV. Usage: python src/eval/batch_eval.py --case KaiB --root results/baseline/fullmsa/KaiB """ from __future__ import annotations import argparse import json import os import re import sys from pathlib import Path # ROOT is only used to render ROOT-relative paths in the output TSV. Default = # repo layout; override with SF_BENCH_ROOT for a relocated benchmark bundle. _ENV_ROOT = os.environ.get("SF_BENCH_ROOT") ROOT = Path(_ENV_ROOT).resolve() if _ENV_ROOT else Path(__file__).resolve().parents[2] # evaluate_prediction.py lives next to this file; import it from there. sys.path.insert(0, str(Path(__file__).resolve().parent)) from evaluate_prediction import evaluate # noqa: E402 PRED_NAME_RE = re.compile( r"^(?P.+?)_unrelaxed_rank_(?P\d+)_alphafold2_ptm_model_(?P\d+)_seed_(?P\d+)\.pdb$" ) def main(argv: list[str] | None = None) -> int: ap = argparse.ArgumentParser() ap.add_argument("--case", required=True, choices=["KaiB", "GA_GB", "Mpt53"]) ap.add_argument("--root", required=True, type=Path, help="Directory containing predicted PDBs (recursive)") ap.add_argument("--out", type=Path, default=None, help="Output TSV; default = {root}/evals.tsv") args = ap.parse_args(argv) pdbs = sorted(args.root.rglob("*_unrelaxed_rank_*.pdb")) if not pdbs: print(f"ERROR: no *_unrelaxed_rank_*.pdb under {args.root}", file=sys.stderr) return 2 out_tsv = args.out or args.root / "evals.tsv" # Dynamic columns: derive from the state keys of the first result sample = evaluate(pdbs[0], args.case) state_keys = list(sample["states"].keys()) cols = [ "subset_id", "model", "seed", "rank", "mean_plddt_overall", "mean_plddt_core", "mean_plddt_switch_3A", "mean_plddt_switch_2A", ] for sk in state_keys: for m in ("rmsd_common_core_A", "rmsd_switch_3A", "rmsd_switch_2A", "tmalign_tm1", "tmalign_tm2", "hit_primary"): cols.append(f"{sk}__{m}") cols.append("pdb") n = len(pdbs) with out_tsv.open("w") as out: out.write("\t".join(cols) + "\n") for i, pdb in enumerate(pdbs, 1): m = PRED_NAME_RE.match(pdb.name) subset = m.group("subset") if m else pdb.stem rank = int(m.group("rank")) if m else -1 model = int(m.group("model")) if m else -1 seed = int(m.group("seed")) if m else -1 try: r = evaluate(pdb, args.case) except Exception as e: print(f"WARN: evaluate failed for {pdb}: {e}", file=sys.stderr) continue # Write JSON sidecar (pdb.with_suffix("").with_name(pdb.stem + "_eval.json")).write_text( json.dumps(r, indent=2, default=float)) # Flatten into TSV row row = [ subset, model, seed, rank, f"{r['mean_plddt_overall']:.2f}", f"{r['mean_plddt_core']:.2f}", f"{r['mean_plddt_switch_3A']:.2f}" if r['mean_plddt_switch_3A'] == r['mean_plddt_switch_3A'] else "NA", f"{r['mean_plddt_switch_2A']:.2f}" if r['mean_plddt_switch_2A'] == r['mean_plddt_switch_2A'] else "NA", ] for sk in state_keys: sv = r["states"].get(sk, {}) for m_ in ("rmsd_common_core_A", "rmsd_switch_3A", "rmsd_switch_2A", "tmalign_tm1", "tmalign_tm2", "hit_primary"): v = sv.get(m_) if v is None or (isinstance(v, float) and v != v): # None or NaN row.append("NA") elif isinstance(v, bool): row.append("1" if v else "0") elif isinstance(v, float): row.append(f"{v:.4f}") else: row.append(str(v)) rel = str(pdb.relative_to(ROOT)) if pdb.is_relative_to(ROOT) else str(pdb) row.append(rel) out.write("\t".join(str(x) for x in row) + "\n") if i % 50 == 0: print(f" {i}/{n} evaluated") try: rel = out_tsv.resolve().relative_to(ROOT) except ValueError: rel = out_tsv print(f"done; {n} predictions -> {rel}") return 0 if __name__ == "__main__": sys.exit(main())