File size: 3,009 Bytes
8d1e644
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
598a072
 
8d1e644
 
 
 
 
 
 
598a072
8d1e644
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Produce the REAL validation numbers — the receipts behind the glass-box claim.

Runs the engine's zero-shot ESM-2 ranking against published deep-mutational-
scanning (DMS) datasets and writes dee/data/benchmarks.json (Spearman ρ +
top-decile precision per dataset). Run this where torch + the ESM weights
exist (the HF Space, or any box with the model); it is the ONLY thing that
should ever populate benchmarks.json — the app never hand-enters numbers.

Input: a manifest JSON, a list of assays:
    [
      {"name": "...", "protein": "GENE_ORG", "sequence": "MSK...",
       "csv": "path/to/dms.csv", "source": "doi:..."},
      ...
    ]
The CSV is ProteinGym-style (a 'mutant'/'mutation' column, a 'DMS_score'/
'score' column). Multi-mutants ('A1C:D5E') are scored as the sum of their
single-site ΔLLs (the same additive assumption the design engine makes — so
this validates exactly what we ship).

Usage:
    python scripts/run_benchmarks.py manifest.json [--model small] [--out dee/data/benchmarks.json]
"""
import argparse
import datetime as dt
import json
import sys
from pathlib import Path

# Make `dee` importable when run from the repo root.
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from dee.core import benchmark as bm          # noqa: E402
from dee.core.dms_seed import parse_proteingym_csv  # noqa: E402


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("manifest")
    ap.add_argument("--model", default="small")
    ap.add_argument("--out", default=str(Path(__file__).resolve().parent.parent / "dee" / "data" / "benchmarks.json"))
    args = ap.parse_args()

    from dee.core import scoring
    scorer = scoring.get_scorer(args.model)

    assays = json.loads(Path(args.manifest).read_text(encoding="utf-8"))
    results = []
    for a in assays:
        recs = parse_proteingym_csv(Path(a["csv"]).read_text(encoding="utf-8"))
        if not recs:
            print(f"  skip {a.get('name')}: no parseable records")
            continue
        labels = [lab for (lab, _v) in recs]
        measured = [v for (_lab, v) in recs]
        scores_df = scorer.score_all_substitutions(a["sequence"])
        predicted = bm.predict_additive(scores_df, labels)
        r = bm.evaluate_dataset(a.get("name", a.get("protein", "?")),
                                a.get("protein", ""), predicted, measured,
                                source=a.get("source", ""))
        results.append(r)
        print(f"  {r.name:24s} n={r.n:6d}  rho={r.spearman}  top10p={r.top_decile_precision}")

    out = {
        "generated_at": dt.datetime.now(dt.timezone.utc).isoformat(),
        "model": args.model,
        "summary": bm.summarize(results),
        "datasets": [r.as_dict() for r in results],
    }
    Path(args.out).write_text(json.dumps(out, indent=2), encoding="utf-8")
    print(f"\nWrote {len(results)} dataset result(s) -> {args.out}")
    print(f"Summary: {out['summary']}")


if __name__ == "__main__":
    main()