"""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)