DooABLe / src /dooable /chemistry.py
pranamanam's picture
Upload 309 files
81ae663 verified
Raw History Blame Contribute Delete
10.7 kB
"""Reaction execution, canonical outcomes, parent metadata, and route replay.
Templates declare graph transformations. Experimental yields and substrate
scope require separate chemical validation. The shipped examples use named
public compounds and are not an Enamine inventory.
"""
from pathlib import Path
import csv, json, re, hashlib
import numpy as np
from rdkit import Chem, DataStructs, RDLogger
from rdkit.Chem import AllChem, Descriptors, QED, rdFingerprintGenerator
from .graph import Graph, Node, Edge
RDLogger.DisableLog("rdApp.warning")
FP = rdFingerprintGenerator.GetMorganGenerator(radius=2, fpSize=128)
def canonical(smiles):
"""Return canonical isomeric SMILES for one valid connected molecule."""
mol = Chem.MolFromSmiles(str(smiles))
if mol is None:
raise ValueError(f"Invalid SMILES: {smiles}")
if len(Chem.GetMolFrags(mol)) != 1:
raise ValueError("A connected molecule is required")
Chem.SanitizeMol(mol)
return Chem.MolToSmiles(mol, isomericSmiles=True)
def features(smiles, depth, budget, terminal=False):
"""Return 131 features comprising 128 Morgan bits and state indicators."""
m = Chem.MolFromSmiles(smiles)
fp = np.zeros(128, dtype=np.float32)
DataStructs.ConvertToNumpyArray(FP.GetFingerprint(m), fp)
return fp.tolist() + [depth / max(budget, 1), float(terminal), 0.0]
def read_catalog(path, limit=None):
"""CSV/TSV input with smiles and id, or native SMILES and Enamine_ID."""
if limit is not None and (int(limit) != limit or limit < 1):
raise ValueError("Catalog row limit must be a positive integer")
p = Path(path)
import gzip
opener = gzip.open if p.suffix == ".gz" else open
rows = []
seen = {}
with opener(p, "rt", newline="") as f:
sample = f.read(8192)
f.seek(0)
if not sample.strip():
raise ValueError("Empty catalog file")
delimiter = "\t" if "\t" in sample.splitlines()[0] else ","
for i, r in enumerate(csv.DictReader(f, delimiter=delimiter)):
r = {k.lower(): v for k, v in r.items()}
smi = canonical(r.get("smiles", ""))
cid = r.get(
"id", r.get("enamine_id", r.get("catalog_id", r.get("code", "")))
)
if not cid:
raise ValueError(f"Missing catalog identifier at row {i+2}")
if cid in seen and seen[cid] != smi:
raise ValueError(f"Conflicting structures for {cid}")
seen[cid] = smi
meta = {k: v for k, v in r.items() if k != "smiles"}
match = re.fullmatch(r"([sm])_(\d+)_([\d_]+)", cid)
if match:
meta.update(
chemistry_class=match[1],
parent_reaction_id=match[2],
parent_reagent_ids=match[3].split("_"),
)
rows.append({"id": cid, "smiles": smi, "metadata": meta})
if limit is not None and len(rows) >= limit:
break
if not rows:
raise ValueError("Empty parent catalog")
# Canonical molecular identity avoids duplicated parent action descriptions.
grouped = {}
for r in rows:
if r["smiles"] in grouped:
grouped[r["smiles"]]["metadata"].setdefault("aliases", []).append(r["id"])
else:
grouped[r["smiles"]] = r
return list(grouped.values())
def reactions(path):
"""Read two-reactant SMARTS templates with unique IDs and finite costs."""
data = json.loads(Path(path).read_text())
items = []
if len({r["id"] for r in data}) != len(data):
raise ValueError("Reaction template IDs must be unique")
for r in data:
rxn = AllChem.ReactionFromSmarts(r["smarts"])
if rxn is None:
raise ValueError("Invalid reaction SMARTS")
if rxn.GetNumReactantTemplates() != 2:
raise ValueError("Templates require parent and one reagent")
cost = float(r.get("cost", 1))
if cost < 0 or not np.isfinite(cost):
raise ValueError("Invalid reaction cost")
items.append((r, rxn))
return items
def products(rxn, parent, reagent):
"""Enumerate unique sanitized products with molecular weight at most 750 Da."""
result = set()
for ps in rxn.RunReactants(
(Chem.MolFromSmiles(parent), Chem.MolFromSmiles(reagent))
):
if len(ps) != 1:
continue
try:
Chem.SanitizeMol(ps[0])
s = canonical(Chem.MolToSmiles(ps[0]))
if Descriptors.MolWt(ps[0]) <= 750:
result.add(s)
except (ValueError, RuntimeError):
continue
return sorted(result)
def build_graph(parents, reagents, templates, budget=2, max_nodes=100000):
"""Expand the supplied inventory through at most budget additional reactions."""
if int(budget) != budget or budget < 0:
raise ValueError("Reaction budget must be a nonnegative integer")
nodes = [Node("root", features=[0.0] * 130 + [1.0])]
edges = []
states = {}
outcomes = {}
def state(smi, k):
key = (smi, k)
if key not in states:
nid = "s" + str(len(states))
states[key] = nid
nodes.append(
Node(
nid,
features=features(smi, k, budget),
metadata={"smiles": smi, "steps": k},
)
)
if len(nodes) > max_nodes:
raise ValueError(
"Graph limit exceeded. Select fewer parents or increase max_nodes explicitly."
)
return states[key]
for p in parents:
sid = state(p["smiles"], 0)
edges.append(
Edge(
"e" + str(len(edges)),
"root",
sid,
0,
{
"kind": "parent",
"parent_id": p["id"],
"smiles": p["smiles"],
"metadata": p.get("metadata", {}),
},
)
)
for k in range(budget + 1):
current = [(s, nid) for (s, depth), nid in states.items() if depth == k]
for smi, sid in current:
if smi not in outcomes:
tid = "t" + str(len(outcomes))
outcomes[smi] = tid
nodes.append(Node(tid, smi, features(smi, 0, budget, True)))
if len(nodes) > max_nodes:
raise ValueError("Graph limit exceeded during terminal creation")
edges.append(
Edge("e" + str(len(edges)), sid, outcomes[smi], 0, {"kind": "stop"})
)
if k == budget:
continue
seen_actions = set()
for spec, rxn in templates:
for reagent in reagents:
for prod in products(rxn, smi, reagent["smiles"]):
if prod == smi:
continue
signature = (spec["id"], reagent["id"], prod)
if signature in seen_actions:
continue
seen_actions.add(signature)
tid = state(prod, k + 1)
edges.append(
Edge(
"e" + str(len(edges)),
sid,
tid,
float(spec.get("cost", 1)),
{
"kind": "reaction",
"template_id": spec["id"],
"smarts": spec["smarts"],
"reagent_id": reagent["id"],
"reagent_smiles": reagent["smiles"],
"input": smi,
"product": prod,
},
)
)
return Graph(
nodes,
edges,
metadata={
"kind": "reaction",
"budget": budget,
"parent_count": len(parents),
"reagent_count": len(reagents),
"template_ids": [x[0]["id"] for x in templates],
"scope": "Complete expansion of the supplied parents, reagents, templates and molecular-weight cap",
"max_molecular_weight": 750,
},
)
def replay(route, budget, parents=None):
"""Re-execute one parent-to-product route under an additional-reaction cap.
Actions must begin with one parent selection and end with one stop. Pass
parsed parent records to check the parent identifier against the inventory.
Malformed records and unsuccessful reaction replay return False.
"""
try:
actions = route["actions"]
if budget < 0 or int(budget) != budget or len(actions) < 2:
return False
if actions[0]["kind"] != "parent" or actions[-1]["kind"] != "stop":
return False
current = canonical(actions[0]["smiles"])
if parents is not None:
inventory = {p["id"]: p["smiles"] for p in parents}
if inventory.get(actions[0]["parent_id"]) != current:
return False
if len(actions) - 2 > budget:
return False
for action in actions[1:-1]:
if action["kind"] != "reaction" or current != action["input"]:
return False
rxn = AllChem.ReactionFromSmarts(action["smarts"])
if rxn is None or rxn.GetNumReactantTemplates() != 2:
return False
if action["product"] not in products(
rxn, current, action["reagent_smiles"]
):
return False
current = action["product"]
return current == canonical(route["outcome"])
except (KeyError, TypeError, ValueError, RuntimeError, IndexError):
return False
def descriptor_rewards(graph, weights=(2.0, 1.0, 1.0)):
"""Return log rewards from QED, logP proximity, and molecular-size utility."""
w = np.asarray(weights, dtype=float)
if w.shape != (3,) or np.any(w < 0) or not np.all(np.isfinite(w)):
raise ValueError("Three nonnegative weights required")
out = {}
for s in graph.terminals:
m = Chem.MolFromSmiles(s)
u = [
QED.qed(m),
np.exp(-(((Descriptors.MolLogP(m) - 2) / 2) ** 2)),
np.exp(-(((Descriptors.MolWt(m) - 300) / 200) ** 2)),
]
out[s] = float(w @ u)
return out