| """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 (`<pdb_stem>_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 |
|
|
| |
| |
| _ENV_ROOT = os.environ.get("SF_BENCH_ROOT") |
| ROOT = Path(_ENV_ROOT).resolve() if _ENV_ROOT else Path(__file__).resolve().parents[2] |
| |
| sys.path.insert(0, str(Path(__file__).resolve().parent)) |
| from evaluate_prediction import evaluate |
|
|
|
|
| PRED_NAME_RE = re.compile( |
| r"^(?P<subset>.+?)_unrelaxed_rank_(?P<rank>\d+)_alphafold2_ptm_model_(?P<model>\d+)_seed_(?P<seed>\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" |
| |
| 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 |
| |
| (pdb.with_suffix("").with_name(pdb.stem + "_eval.json")).write_text( |
| json.dumps(r, indent=2, default=float)) |
| |
| 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): |
| 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()) |
|
|