File size: 3,995 Bytes
81ae663
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Summarize generated candidates and recorded routes using a common schema."""

import argparse
import json
from pathlib import Path
import numpy as np
import pandas as pd
from rdkit import Chem, DataStructs
from rdkit.Chem import QED, rdFingerprintGenerator
from dooable.chemistry import canonical, replay
from dooable.metrics import hypervolume_2d


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--samples", required=True)
    p.add_argument("--scores", required=True)
    p.add_argument("--budget", required=True, type=int)
    p.add_argument("--requested", required=True, type=int)
    p.add_argument("--oracle-calls", required=True, type=int)
    p.add_argument("--method", required=True)
    p.add_argument("--seed", type=int, default=0)
    p.add_argument("--target", default="BACE_public")
    p.add_argument("--output", required=True)
    args = p.parse_args()
    rows = [
        json.loads(x) for x in Path(args.samples).read_text().splitlines() if x.strip()
    ]
    if args.requested < len(rows) or args.requested < 1:
        raise ValueError(
            "Requested count must be positive and at least the output count"
        )
    scores = pd.read_csv(args.scores)
    scores["smiles"] = scores.smiles.map(canonical)
    if scores.smiles.duplicated().any():
        raise ValueError("Score rows must be unique by canonical structure")
    required = [
        "smiles",
        "bace_utility",
        "caco2_utility",
        "predicted_bace_pIC50",
        "predicted_caco2_log10_cm_s",
    ]
    if any(x not in scores for x in required):
        raise ValueError("Score CSV must contain the public property-score schema")
    feasible = [r for r in rows if replay(r, args.budget)]
    molecules = sorted({canonical(r["outcome"]) for r in feasible})
    measured = scores.set_index("smiles").reindex(molecules)
    if measured[required[1:]].isna().any().any():
        raise ValueError("Every feasible outcome requires finite recorded scores")
    measured["reward"] = 0.5 * (measured.bace_utility + measured.caco2_utility)
    top = measured.sort_values("reward", ascending=False).head(100)
    generator = rdFingerprintGenerator.GetMorganGenerator(radius=2, fpSize=1024)
    fps = [generator.GetFingerprint(Chem.MolFromSmiles(s)) for s in top.index]
    distances = [
        1 - DataStructs.TanimotoSimilarity(fps[i], fps[j])
        for i in range(len(fps))
        for j in range(i)
    ]
    result = {
        "method": args.method,
        "target": args.target,
        "seed": args.seed,
        "budget": args.budget,
        "requested": args.requested,
        "returned": len(rows),
        "oracle_calls": args.oracle_calls,
        "replay_fraction": len(feasible) / args.requested,
        "unique_feasible": len(molecules),
        "top_count": len(top),
        "mean_reaction_steps": (
            float(
                np.mean(
                    [
                        sum(a.get("kind") == "reaction" for a in r["actions"])
                        for r in feasible
                    ]
                )
            )
            if feasible
            else None
        ),
        "top_mean_bace_pIC50": (
            float(top.predicted_bace_pIC50.mean()) if len(top) else None
        ),
        "top_mean_caco2_log10_cm_s": (
            float(top.predicted_caco2_log10_cm_s.mean()) if len(top) else None
        ),
        "top_mean_qed": (
            float(np.mean([QED.qed(Chem.MolFromSmiles(s)) for s in top.index]))
            if len(top)
            else None
        ),
        "top_internal_diversity": float(np.mean(distances)) if distances else None,
        "hypervolume": hypervolume_2d(
            measured[["bace_utility", "caco2_utility"]].to_numpy()
        ),
    }
    out = Path(args.output)
    out.parent.mkdir(parents=True, exist_ok=True)
    out.write_text(json.dumps(result, indent=2, allow_nan=False))
    print(json.dumps(result, allow_nan=False))


if __name__ == "__main__":
    main()