Download eval/batch_eval.py from ChatterjeeLab/SF-Cluster: direct link, hf CLI and curl.
- Browser
- Download file 4.63 kB
-
https://huggingface.co/ChatterjeeLab/SF-Cluster/resolve/main/eval/batch_eval.py
- Command line
-
hf download hf://ChatterjeeLab/SF-Cluster/eval/batch_eval.py
-
curl -L -o batch_eval.py https://huggingface.co/ChatterjeeLab/SF-Cluster/resolve/main/eval/batch_eval.py
4.63 kB
| """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 | |
| # 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<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" | |
| # 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()) | |