SF-Cluster / eval /batch_eval.py
chq1155's picture
Add benchmark reproduction: CPU scoring/eval (evaluate_prediction+batch_eval), reference structures, region manifests, headline prediction sets, reproduce_benchmark.py (reproduces main minority_hit_rate table; 15/15 cells verified by re-scoring 1440 PDBs)
f4e8048 verified
Raw
History Blame Contribute Delete
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())