File size: 4,632 Bytes
f4e8048
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
"""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())