| |
| """main_benchmark §9.1 evaluator. |
| |
| For each prediction PDB under <root>, computes: |
| - mean_plddt_overall, mean_plddt_switch, mean_plddt_core |
| - For each state (a, b): Cα RMSD common_core, TM-score, hit_primary (§9.1 binding 3-condition) |
| |
| Regions come from data/main_benchmark/annotations/<case_id>/state_region_FINAL.tsv |
| Reference PDBs come from data/main_benchmark/structures/<case_id>/{state_a,state_b}.pdb |
| |
| ID-pattern cases (state_a == state_b): single hit column hit_primary_state_a. |
| Per IDP_EVAL_OVERRIDE: treat the whole query region as the target region; |
| RMSD ≤3Å AND mean_pLDDT ≥70 → hit_primary=True. switch_region pLDDT check is |
| skipped because every residue is marked switch_region. |
| |
| Usage: |
| python scripts/main_benchmark_evaluate.py --case <case_id> --root <pred_dir> |
| """ |
| from __future__ import annotations |
| import argparse |
| import csv |
| import json |
| import os |
| import re |
| import shutil |
| import subprocess |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import yaml |
| from Bio.PDB import PDBParser, Superimposer |
|
|
| |
| |
| DATA_ROOT = Path(os.environ.get( |
| "SF_MAIN_BENCH_DATA", Path(__file__).resolve().parents[1] / "bench" / "main_benchmark")) |
| |
| |
| TMALIGN = shutil.which("TMalign") |
|
|
| PRED_NAME_RE = re.compile( |
| r"^(?P<subset>.+?)_unrelaxed_rank_(?P<rank>\d+)_alphafold2_ptm_model_(?P<model>\d+)_seed_(?P<seed>\d+)\.pdb$" |
| ) |
|
|
| _CASES_CACHE: dict[str, dict] | None = None |
|
|
|
|
| def load_case(case_id: str) -> dict: |
| global _CASES_CACHE |
| if _CASES_CACHE is None: |
| cases = yaml.safe_load((DATA_ROOT / "cases.yaml").read_text())["cases"] |
| _CASES_CACHE = {c["case_id"]: c for c in cases} |
| return _CASES_CACHE[case_id] |
|
|
|
|
| def load_regions(case_id: str) -> dict: |
| path = DATA_ROOT / "annotations" / case_id / "state_region_FINAL.tsv" |
| common_core: list[int] = [] |
| switch_3a: list[int] = [] |
| state_a_resnum: dict[int, int] = {} |
| with path.open() as f: |
| for line in f: |
| if line.startswith("#"): |
| continue |
| line = line.rstrip("\n") |
| if not line: |
| continue |
| parts = line.split("\t") |
| if parts[0] == "residue_index_query": |
| continue |
| qi = int(parts[0]) |
| try: |
| resnum_a = int(parts[1]) |
| except ValueError: |
| resnum_a = qi |
| cc = int(parts[3]) |
| sr = int(parts[4]) |
| if cc: |
| common_core.append(qi) |
| if sr: |
| switch_3a.append(qi) |
| state_a_resnum[qi] = resnum_a |
| return {"common_core": common_core, "switch_3A": switch_3a, |
| "state_a_resnum": state_a_resnum} |
|
|
|
|
| def is_id_case(case_id: str) -> bool: |
| return case_id.startswith("SFB_ID_") |
|
|
|
|
| def read_ca_per_resid(pdb: Path, chain: str | None = None) -> dict[int, np.ndarray]: |
| s = PDBParser(QUIET=True).get_structure("x", str(pdb)) |
| m = next(iter(s)) |
| for c in m: |
| if chain is not None and c.id != chain: |
| continue |
| out: dict[int, np.ndarray] = {} |
| for r in c: |
| if r.id[0] != " ": |
| continue |
| if "CA" not in r: |
| continue |
| out[r.id[1]] = r["CA"].coord |
| if out: |
| return out |
| raise RuntimeError(f"no Cα in {pdb} chain={chain}") |
|
|
|
|
| def read_ca_bfactors(pdb: Path, chain: str | None = None) -> dict[int, float]: |
| s = PDBParser(QUIET=True).get_structure("x", str(pdb)) |
| m = next(iter(s)) |
| for c in m: |
| if chain is not None and c.id != chain: |
| continue |
| out: dict[int, float] = {} |
| for r in c: |
| if r.id[0] != " ": |
| continue |
| if "CA" not in r: |
| continue |
| out[r.id[1]] = float(r["CA"].bfactor) |
| if out: |
| return out |
| raise RuntimeError(f"no Cα B-factors in {pdb} chain={chain}") |
|
|
|
|
| def rmsd_on_residues(pred_ca: dict[int, np.ndarray], |
| ref_ca: dict[int, np.ndarray], |
| pred_to_ref: dict[int, int], |
| residues: list[int]) -> tuple[float, int]: |
| pairs = [(qi, pred_to_ref[qi]) for qi in residues |
| if qi in pred_ca and qi in pred_to_ref and pred_to_ref[qi] in ref_ca] |
| if len(pairs) < 3: |
| return float("nan"), len(pairs) |
| from Bio.PDB.Atom import Atom |
| pred_atoms = [Atom("CA", pred_ca[qi], 1.0, 1.0, " ", "CA", 1, "C") for qi, _ in pairs] |
| ref_atoms = [Atom("CA", ref_ca[ri], 1.0, 1.0, " ", "CA", 1, "C") for _, ri in pairs] |
| sup = Superimposer() |
| sup.set_atoms(ref_atoms, pred_atoms) |
| return float(sup.rms), len(pairs) |
|
|
|
|
| def tmalign(pdb_a: Path, pdb_b: Path) -> tuple[float, float, int]: |
| if not TMALIGN: |
| return float("nan"), float("nan"), 0 |
| try: |
| p = subprocess.run([TMALIGN, str(pdb_a), str(pdb_b)], |
| capture_output=True, text=True, timeout=120) |
| except Exception: |
| return float("nan"), float("nan"), 0 |
| t1 = t2 = float("nan"); aln = 0 |
| for line in p.stdout.splitlines(): |
| if line.startswith("TM-score=") and "normalized by length of Chain_1" in line: |
| try: t1 = float(line.split()[1]) |
| except: pass |
| elif line.startswith("TM-score=") and "normalized by length of Chain_2" in line: |
| try: t2 = float(line.split()[1]) |
| except: pass |
| elif line.startswith("Aligned length="): |
| parts = line.replace(",", " ").split() |
| for i, q in enumerate(parts): |
| if q == "length=": |
| try: aln = int(parts[i + 1]) |
| except: pass |
| return t1, t2, aln |
|
|
|
|
| def get_state_chain(case: dict, which: str) -> str: |
| """Get the chain id within state PDB. The structures/ files use the curated chain |
| from `case_definitions.py` build_v2 — they are usually chain A in the file. |
| Try chain A first; fall back to the chain id from cases.yaml.""" |
| return case[f"state_{which}_chain"] |
|
|
|
|
| def state_chain_in_file(pdb: Path) -> str | None: |
| """Return first chain id present in the PDB file (with Cα atoms).""" |
| s = PDBParser(QUIET=True).get_structure("x", str(pdb)) |
| m = next(iter(s)) |
| for c in m: |
| for r in c: |
| if r.id[0] == " " and "CA" in r: |
| return c.id |
| return None |
|
|
|
|
| def evaluate_pdb(pdb: Path, case_id: str) -> dict: |
| case = load_case(case_id) |
| regions = load_regions(case_id) |
| id_case = is_id_case(case_id) |
|
|
| pred_ca = read_ca_per_resid(pdb, chain="A") |
| pred_plddt = read_ca_bfactors(pdb, chain="A") |
| |
| mean_plddt_overall = float(np.mean(list(pred_plddt.values()))) |
| |
| core_p = [pred_plddt[r] for r in regions["common_core"] if r in pred_plddt] |
| sw_p = [pred_plddt[r] for r in regions["switch_3A"] if r in pred_plddt] |
| mean_plddt_core = float(np.mean(core_p)) if core_p else float("nan") |
| mean_plddt_switch = float(np.mean(sw_p)) if sw_p else float("nan") |
|
|
| result = { |
| "pdb": str(pdb), |
| "case_id": case_id, |
| "id_case": id_case, |
| "pred_len": len(pred_ca), |
| "mean_plddt_overall": mean_plddt_overall, |
| "mean_plddt_core": mean_plddt_core, |
| "mean_plddt_switch_3A": mean_plddt_switch, |
| "states": {}, |
| } |
|
|
| |
| |
| |
| |
| |
| pred_to_state_a = dict(regions["state_a_resnum"]) |
|
|
| for which in ["a", "b"]: |
| pdb_state = DATA_ROOT / "structures" / case_id / f"state_{which}.pdb" |
| if not pdb_state.exists(): |
| continue |
| ch = state_chain_in_file(pdb_state) |
| ref_ca = read_ca_per_resid(pdb_state, chain=ch) |
| |
| if which == "a": |
| pred_to_ref = pred_to_state_a |
| else: |
| |
| |
| |
| |
| |
| ref_keys = set(ref_ca.keys()) |
| num_overlap = sum(1 for rn in pred_to_state_a.values() if rn in ref_keys) |
| if pred_to_state_a and num_overlap >= 0.5 * len(pred_to_state_a): |
| pred_to_ref = pred_to_state_a |
| else: |
| ref_sorted = sorted(ref_ca.keys()) |
| pred_sorted = sorted(pred_ca.keys()) |
| pred_to_ref = {p: r for p, r in zip(pred_sorted, ref_sorted)} |
|
|
| rms_core, n_core = rmsd_on_residues(pred_ca, ref_ca, pred_to_ref, |
| regions["common_core"]) |
| rms_sw, n_sw = rmsd_on_residues(pred_ca, ref_ca, pred_to_ref, |
| regions["switch_3A"]) |
| t1, t2, aln = tmalign(pdb, pdb_state) |
|
|
| if id_case: |
| |
| rms_target = rms_sw if not np.isnan(rms_sw) else rms_core |
| hit = (not np.isnan(rms_target) and rms_target <= 3.0 |
| and mean_plddt_overall >= 70.0) |
| else: |
| |
| hit = (not np.isnan(rms_core) and rms_core <= 3.0 |
| and mean_plddt_overall >= 70.0 |
| and not np.isnan(mean_plddt_switch) |
| and mean_plddt_switch >= 70.0) |
|
|
| result["states"][f"state_{which}"] = { |
| "rmsd_common_core_A": rms_core, |
| "rmsd_switch_3A": rms_sw, |
| "n_common_core_aligned": n_core, |
| "n_switch_3A_aligned": n_sw, |
| "tmalign_tm1": t1, |
| "tmalign_tm2": t2, |
| "tmalign_aligned_len": aln, |
| "hit_primary": hit, |
| } |
|
|
| if id_case: |
| |
| break |
|
|
| return result |
|
|
|
|
| def main(argv: list[str] | None = None) -> int: |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--case", required=True) |
| ap.add_argument("--root", required=True, type=Path) |
| ap.add_argument("--out", type=Path, default=None) |
| args = ap.parse_args(argv) |
|
|
| pdbs = sorted(args.root.rglob("*_unrelaxed_rank_*.pdb")) |
| if not pdbs: |
| print(f"ERROR: no PDBs under {args.root}", file=sys.stderr) |
| return 2 |
| out_tsv = args.out or args.root / "evals.tsv" |
| sample = evaluate_pdb(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"] |
| for sk in state_keys: |
| for m in ("rmsd_common_core_A", "rmsd_switch_3A", |
| "tmalign_tm1", "tmalign_tm2", "hit_primary"): |
| cols.append(f"{sk}__{m}") |
| cols.append("pdb") |
|
|
| with out_tsv.open("w") as out: |
| out.write("\t".join(cols) + "\n") |
| for pdb in pdbs: |
| m = PRED_NAME_RE.match(pdb.name) |
| sid = 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(pdb, args.case) |
| except Exception as e: |
| print(f"WARN: eval failed {pdb}: {e}", file=sys.stderr) |
| continue |
| (pdb.with_name(pdb.stem + "_eval.json")).write_text( |
| json.dumps(r, indent=2, default=float)) |
| row = [sid, model, seed, rank, |
| f"{r['mean_plddt_overall']:.2f}", |
| f"{r['mean_plddt_core']:.2f}" if r['mean_plddt_core'] == r['mean_plddt_core'] else "NA", |
| f"{r['mean_plddt_switch_3A']:.2f}" if r['mean_plddt_switch_3A'] == r['mean_plddt_switch_3A'] else "NA"] |
| for sk in state_keys: |
| sv = r["states"].get(sk, {}) |
| for k in ("rmsd_common_core_A", "rmsd_switch_3A", |
| "tmalign_tm1", "tmalign_tm2", "hit_primary"): |
| v = sv.get(k) |
| 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)) |
| row.append(str(pdb)) |
| out.write("\t".join(str(x) for x in row) + "\n") |
| print(f"wrote {len(pdbs)} preds -> {out_tsv}") |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| sys.exit(main()) |
|
|