Commit df442c57 authored by Bharath Ramsundar's avatar Bharath Ramsundar
Browse files

Merge pull request #157 from joegomes/mwsplit

Molecular Weight Splitter Class
parents 53c02d0b 0d2a7d50
Loading
Loading
Loading
Loading
+31 −0
Original line number Diff line number Diff line
@@ -95,6 +95,37 @@ class Splitter(object):
    """
    raise NotImplementedError

class MolecularWeightSplitter(Splitter):
  """
  Class for doing data splits by molecular weight.
  """
  def split(self, samples, seed=None, frac_train=.8, frac_valid=.1,
            frac_test=.1, log_every_n=None):
    """
    Splits internal compounds into train/validation/test using the MW calculated
    by SMILES string.
    """

    np.testing.assert_almost_equal(frac_train + frac_valid + frac_test, 1.)
    np.random.seed(seed)

    df = samples.compounds_df
    mws = []
    for _, row in df.iterrows():
        mol = Chem.MolFromSmiles(row['smiles'])
        mw = Chem.rdMolDescriptors.CalcExactMolWt(mol)
        mws.append(mw)

    # Sort by increasing MW
    mws = np.array(mws)
    sortidx = np.argsort(mws)

    train_cutoff = frac_train * len(sortidx)
    valid_cutoff = (frac_train+frac_valid) * len(sortidx)

    return (sortidx[:train_cutoff], sortidx[train_cutoff:valid_cutoff],
            sortidx[valid_cutoff:])

class RandomSplitter(Splitter):
  """
  Class for doing random data splits.