DooABLe / tests /test_workflow.py
pranamanam's picture
Upload 309 files
81ae663 verified
Raw
History Blame Contribute Delete
3.34 kB
import json
from pathlib import Path
import numpy as np
import pytest
from dooable.chemistry import (
read_catalog,
reactions,
build_graph,
replay,
descriptor_rewards,
)
from dooable.graph import Graph, toy_graph
from dooable.exact import solve, sample, endpoint_distribution
from dooable.learning import train, load_model
from dooable.properties import scaffold_split
ROOT = Path(__file__).resolve().parents[1]
def test_chemical_route_replay_and_serialization(tmp_path):
ex = ROOT / "data/examples"
g = build_graph(
read_catalog(ex / "parents.csv"),
read_catalog(ex / "reagents.csv"),
reactions(ex / "reactions.json"),
2,
)
g.save(tmp_path / "graph.json")
gg = Graph.load(tmp_path / "graph.json")
assert len(gg.nodes) == len(g.nodes)
sol = solve(gg, descriptor_rewards(gg), 0.7)
rows = sample(gg, sol.forward, 200, 7)
assert all(replay(r, 2) for r in rows)
for y, v in gg.terminals.items():
assert gg.nodes[v].outcome == y
def test_neural_training_and_checkpoint(tmp_path):
g = toy_graph(4)
m, h = train(
g, {"A": 0.0, "B": 0.0}, steps=600, batch_size=32, seed=11, output=tmp_path
)
loaded, _ = load_model(tmp_path)
np.testing.assert_allclose(m.probabilities(), loaded.probabilities(), atol=1e-8)
assert h[-1]["endpoint_tv"] < 0.06
def test_scaffold_split_separation():
from rdkit.Chem.Scaffolds import MurckoScaffold
smi = [
"CC",
"CCC",
"CCO",
"c1ccccc1",
"Cc1ccccc1",
"c1ccncc1",
"C1CCCCC1",
"CC1CCCCC1",
"c1ccoc1",
]
tr, te = scaffold_split(smi, 3, 0.7)
scaff = lambda inds: {
MurckoScaffold.MurckoScaffoldSmiles(smiles=smi[i]) for i in inds
}
assert not scaff(tr) & scaff(te)
def test_catalog_conflicting_identifiers(tmp_path):
f = tmp_path / "parents.csv"
f.write_text("id,smiles\na,CC\na,CCC\n")
with pytest.raises(ValueError, match="Conflicting"):
read_catalog(f)
def test_resumed_training_matches_uninterrupted(tmp_path):
g = toy_graph(3)
rewards = {"A": 0.0, "B": 0.0}
full, _ = train(g, rewards, steps=30, seed=2, batch_size=16)
train(g, rewards, steps=12, seed=2, batch_size=16, output=tmp_path)
resumed, _ = train(g, rewards, steps=30, seed=2, batch_size=16, resume=tmp_path)
np.testing.assert_allclose(full.probabilities(), resumed.probabilities(), atol=1e-9)
def test_replay_requires_parent_and_stop():
assert replay(
{
"actions": [{"kind": "parent", "smiles": "CC"}, {"kind": "stop"}],
"outcome": "CC",
},
0,
)
assert not replay(
{"actions": [{"kind": "parent", "smiles": "CC"}], "outcome": "CC"}, 0
)
assert not replay(
{
"actions": [
{"kind": "parent", "smiles": "CC"},
{"kind": "parent", "smiles": "CC"},
{"kind": "stop"},
],
"outcome": "CC",
},
2,
)
def test_terminal_merging_ablation_changes_endpoint_mass():
from dooable.ablations import duplicate_endpoint_policy
g = toy_graph(5)
r = {"A": 0.0, "B": 0.0}
p = endpoint_distribution(g, duplicate_endpoint_policy(g, r, 0.7))
assert p["A"] == pytest.approx(5 / 6)