Download scripts/evaluate_candidates.py from ChatterjeeLab/DooABLe: direct link, hf CLI and curl.
- Browser
- Download file 4 kB
-
https://huggingface.co/ChatterjeeLab/DooABLe/resolve/main/scripts/evaluate_candidates.py
- Command line
-
hf download hf://ChatterjeeLab/DooABLe/scripts/evaluate_candidates.py
-
curl -L -o evaluate_candidates.py https://huggingface.co/ChatterjeeLab/DooABLe/resolve/main/scripts/evaluate_candidates.py
4 kB
| """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() | |