ReMEDi / src /remedi /splits.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw History Blame Contribute Delete
2.03 kB
"""Molecule-disjoint or scaffold-disjoint splits before feature fitting."""
import numpy as np
import pandas as pd
from .chemistry import canonical_smiles, molecule_id
def make_splits(frame, seed=0, fractions=(.6, .15, .1, .15), scaffold=False):
if len(fractions) != 4 or not np.isclose(sum(fractions), 1) or min(fractions) <= 0:
raise ValueError("Four positive split fractions must sum to one")
frame = frame[["smiles"]].drop_duplicates().copy()
frame["smiles"] = frame.smiles.map(canonical_smiles)
frame = frame.drop_duplicates("smiles")
frame["molecule_id"] = frame.smiles.map(molecule_id)
frame["group"] = frame.molecule_id
if scaffold:
from rdkit import Chem
from rdkit.Chem.Scaffolds import MurckoScaffold
def key(s):
scaffold = MurckoScaffold.MurckoScaffoldSmiles(mol=Chem.MolFromSmiles(s))
# All acyclic structures remain one group in this conservative implementation.
return scaffold or "ACYCLIC"
frame["group"] = frame.smiles.map(key)
groups = frame.groupby("group").size()
if len(groups) < 8:
raise ValueError("At least eight independent molecule/scaffold groups are required")
rng = np.random.default_rng(seed)
order = groups.index.to_numpy()[rng.permutation(len(groups))]
# Allocate whole groups. Each split receives at least one group.
counts = np.maximum(1, np.floor(np.asarray(fractions) * len(order)).astype(int))
while counts.sum() < len(order):
counts[np.argmax(np.asarray(fractions) * len(order) - counts)] += 1
while counts.sum() > len(order):
eligible = np.where(counts > 1, counts - np.asarray(fractions)*len(order), -np.inf)
counts[np.argmax(eligible)] -= 1
assignments = {}
pos = 0
for name, n in zip(["train", "tune", "calibration", "test"], counts):
assignments.update({g: name for g in order[pos:pos+n]})
pos += n
frame["split"] = frame.group.map(assignments)
return frame.reset_index(drop=True)