File size: 2,030 Bytes
3f98d52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
"""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)