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

Merge pull request #151 from rbharath/generalized_splitter

Generalized Splits
parents 7fc142d4 0090a524
Loading
Loading
Loading
Loading
+0 −19
Original line number Diff line number Diff line
@@ -275,25 +275,6 @@ class Dataset(object):
    df = self.metadata_df
    update_mean_and_std(df)
 
  def balance_positives_and_negatives(self):
    """For binary datasets, balance pos and neg examples."""
    labels = self.get_labels()
    # Ensure dataset is binary
    np.testing.assert_allclose(sorted(np.unique(labels)), np.array([0., 1.]))
    weights = []
    # TODO(rbharath): This doesn't deal with zeroed out labels.
    for ind, task in enumerate(self.get_task_names()):
      task_labels = labels[:, ind]
      num_positives = np.count_nonzero(task_labels)
      num_negatives = len(task_labels) - num_positives
      if num_positives > 0:
        pos_weight = float(num_negatives)/num_positives
      else:
        pos_weight = 1
      neg_weight = 1
      weights.append((pos_weight, neg_weight))
    return weights
    

def compute_sums_and_nb_sample(tensor, W=None):
  """
+6 −25
Original line number Diff line number Diff line
@@ -18,24 +18,12 @@ from deepchem.datasets import Dataset
from deepchem.featurizers.featurize import DataFeaturizer
from deepchem.featurizers.fingerprints import CircularFingerprint
from deepchem.transformers import NormalizationTransformer
from deepchem.splits.tests import TestSplitAPI

class TestDatasetAPI(unittest.TestCase):
class TestDatasetAPI(TestSplitAPI):
  """
  Shared API for testing with dataset objects. 
  """
  def setUp(self):
    self.current_dir = os.path.dirname(os.path.abspath(__file__))
    self.test_data_dir = os.path.join(self.current_dir, "../../models/test")
    self.smiles_field = "smiles"
    self.feature_dir = tempfile.mkdtemp()
    self.samples_dir = tempfile.mkdtemp()
    self.data_dir = tempfile.mkdtemp()

  def tearDown(self):
    shutil.rmtree(self.feature_dir)
    shutil.rmtree(self.samples_dir)
    shutil.rmtree(self.data_dir)

  # TODO(rbharath): There should be a more natural way to create a dataset
  # object, perhaps just starting from (Xs, ys, ws)
  def _create_dataset(self, compound_featurizers, complex_featurizers,
@@ -45,21 +33,15 @@ class TestDatasetAPI(unittest.TestCase):
                      user_specified_features=None,
                      split_field=None,
                      shard_size=100):
    # Featurize input
    featurizers = compound_featurizers + complex_featurizers

    input_file = os.path.join(self.test_data_dir, input_file)
    featurizer = DataFeaturizer(tasks=tasks,
                                smiles_field=self.smiles_field,
    samples = self._gen_samples(
        compound_featurizers, complex_featurizers,
        input_transformer_classes, output_transformer_classes,
        input_file, tasks,
        protein_pdb_field=protein_pdb_field,
        ligand_pdb_field=ligand_pdb_field,
                                compound_featurizers=compound_featurizers,
                                complex_featurizers=complex_featurizers,
        user_specified_features=user_specified_features,
        split_field=split_field,
                                verbosity="low")

    samples = featurizer.featurize(input_file, self.feature_dir, self.samples_dir,
        shard_size=shard_size)
    use_user_specified_features = (user_specified_features is not None)
    dataset = Dataset(data_dir=self.data_dir, samples=samples, 
@@ -93,7 +75,6 @@ class TestDatasetAPI(unittest.TestCase):
        input_transformer_classes, output_transformer_classes,
        input_file, task_types.keys())


  def _load_multitask_data(self):
    """Load example multitask data."""
    compound_featurizers = [CircularFingerprint(size=1024)]
+0 −12
Original line number Diff line number Diff line
@@ -87,15 +87,3 @@ class TestBasicDatasetAPI(TestDatasetAPI):
    np.testing.assert_allclose(comp_y_means, y_means)
    np.testing.assert_allclose(comp_X_stds, X_stds)
    np.testing.assert_allclose(comp_y_stds, y_stds)

  def test_balance_positives_and_negatives(self):
    """Test balancing of positive and negative examples."""
    multitask_dataset = self._load_multitask_data()
    weights = multitask_dataset.balance_positives_and_negatives()
    X, y, w, ids = multitask_dataset.to_numpy()
    for ind, task in enumerate(multitask_dataset.get_task_names()):
      task_labels = y[:, ind]
      num_positives = np.count_nonzero(task_labels)
      num_negatives = len(task_labels) - num_positives
      pos_weight, neg_weight = weights[ind]
      assert np.isclose(num_positives * pos_weight, num_negatives * neg_weight)
+0 −134
Original line number Diff line number Diff line
@@ -17,19 +17,11 @@ from deepchem.utils.save import log
from deepchem.utils.save import save_to_disk
from deepchem.utils.save import load_from_disk
from deepchem.utils.save import load_pandas_from_disk
from deepchem.utils import ScaffoldGenerator
from deepchem.featurizers.nnscore import NNScoreComplexFeaturizer
import multiprocessing as mp
from functools import partial
import dill

def generate_scaffold(smiles, include_chirality=False):
  """Compute the Bemis-Murcko scaffold for a SMILES string."""
  mol = Chem.MolFromSmiles(smiles)
  engine = ScaffoldGenerator(include_chirality=include_chirality)
  scaffold = engine.get_scaffold(mol)
  return scaffold

def _check_validity(compounds_df):
  """Ensure that columns of compound_df contain required elements."""
  if not set(FeaturizedSamples.colnames).issubset(compounds_df.keys()):
@@ -417,129 +409,3 @@ class FeaturizedSamples(object):
        if row["mol_id"] in compound_ids:
          visible_inds.append(ind)
      yield df.loc[visible_inds]

  def train_valid_test_split(self, splittype, train_dir=None,
                             valid_dir=None, test_dir=None, frac_train=.8,
                             frac_valid=.1, frac_test=.1, seed=None,
                             log_every_n=1000, reload=False):
    """
    Splits self into train/validation/test sets.

    Returns FeaturizedDataset objects.
    """
    if not reload:
      if splittype == "random":
        train_inds, valid_inds, test_inds = self._random_split(
            seed=seed, frac_train=frac_train, frac_test=frac_test,
            frac_valid=frac_valid)
      elif splittype == "scaffold":
        train_inds, valid_inds, test_inds = self._scaffold_split(
            frac_train=frac_train, frac_test=frac_test,
            frac_valid=frac_valid, log_every_n=log_every_n)
      elif splittype == "specified":
        train_inds, valid_inds, test_inds = self._specified_split()
      else:
        raise ValueError("improper splittype.")
    train_samples, valid_samples, test_samples = None, None, None
    dataset_files = self.dataset_files
    if train_dir is not None:
      train_samples = FeaturizedSamples(samples_dir=train_dir, 
                                        dataset_files=dataset_files,
                                        featurizers=self.featurizers,
                                        verbosity=self.verbosity,
                                        reload=False)
      if not reload:
        train_samples._set_compound_df(self.compounds_df.iloc[train_inds])
    if test_dir is not None:
      test_samples = FeaturizedSamples(samples_dir=test_dir, 
                                       dataset_files=dataset_files,
                                       featurizers=self.featurizers,
                                       verbosity=self.verbosity,
                                       reload=False)
      if not reload:
        test_samples._set_compound_df(self.compounds_df.iloc[test_inds])
    if valid_dir is not None:
      valid_samples = FeaturizedSamples(samples_dir=valid_dir, 
                                       dataset_files=dataset_files,
                                       featurizers=self.featurizers,
                                       verbosity=self.verbosity,
                                       reload=False)
      if not reload:
        valid_samples._set_compound_df(self.compounds_df.iloc[valid_inds])

    return train_samples, valid_samples, test_samples

  def train_test_split(self, splittype, train_dir, test_dir, seed=None,
                       frac_train=.8, reload=False):
    """
    Splits self into train/test sets.

    Returns FeaturizedDataset objects.
    """
    train_samples, _, test_samples = self.train_valid_test_split(
        splittype, train_dir, valid_dir=None, test_dir=test_dir,
        frac_train=frac_train, frac_test=1-frac_train, frac_valid=0.,
        reload=False)
    return train_samples, test_samples

  def _random_split(self, seed=None, frac_train=.8, frac_valid=.1,
                    frac_test=.1):
    """
    Splits internal compounds randomly into train/validation/test.
    """
    np.testing.assert_almost_equal(frac_train + frac_valid + frac_test, 1.)
    np.random.seed(seed)
    train_cutoff = frac_train * len(self.compounds_df)
    valid_cutoff = (frac_train+frac_valid) * len(self.compounds_df)
    shuffled = np.random.permutation(range(len(self.compounds_df)))
    return (shuffled[:train_cutoff], shuffled[train_cutoff:valid_cutoff],
            shuffled[valid_cutoff:])

  def _scaffold_split(self, frac_train=.8, frac_valid=.1, frac_test=.1, log_every_n=1000):
    """
    Splits internal compounds into train/validation/test by scaffold.
    """
    np.testing.assert_almost_equal(frac_train + frac_valid + frac_test, 1.)
    scaffolds = {}
    log("About to generate scaffolds", self.verbosity)
    for ind, row in self.compounds_df.iterrows():
      if self.verbosity is not None and ind % log_every_n == 0:
        log("Generating scaffold %d/%d" % (ind, len(self.compounds_df)),
            self.verbosity)
      scaffold = generate_scaffold(row["smiles"])
      if scaffold not in scaffolds:
        scaffolds[scaffold] = [ind]
      else:
        scaffolds[scaffold].append(ind)
    # Sort from largest to smallest scaffold sets
    scaffold_sets = [scaffold_set for (scaffold, scaffold_set) in
                     sorted(scaffolds.items(), key=lambda x: -len(x[1]))]
    train_cutoff = frac_train * len(self.compounds_df)
    valid_cutoff = (frac_train+frac_valid) * len(self.compounds_df)
    train_inds, valid_inds, test_inds = [], [], []
    log("About to sort in scaffold sets", self.verbosity)
    for scaffold_set in scaffold_sets:
      if len(train_inds) + len(scaffold_set) > train_cutoff:
        if len(train_inds) + len(valid_inds) + len(scaffold_set) > valid_cutoff:
          test_inds += scaffold_set
        else:
          valid_inds += scaffold_set
      else:
        train_inds += scaffold_set
    return train_inds, valid_inds, test_inds

  def _specified_split(self):
    """
    Splits internal compounds into train/validation/test by user-specification.
    """
    train_inds, valid_inds, test_inds = [], [], []
    for ind, row in self.compounds_df.iterrows():
      if row["split"].lower() == "train":
        train_inds.append(ind)
      elif row["split"].lower() in ["valid", "validation"]:
        valid_inds.append(ind)
      elif row["split"].lower() == "test":
        test_inds.append(ind)
      else:
        raise ValueError("Missing required split information.")
    return train_inds, valid_inds, test_inds
+14 −4
Original line number Diff line number Diff line
@@ -13,6 +13,9 @@ import os
import unittest
import tempfile
import shutil
from deepchem.splits import RandomSplitter
from deepchem.splits import ScaffoldSplitter
from deepchem.splits import SpecifiedSplitter
from deepchem.featurizers.featurize import DataFeaturizer
from deepchem.featurizers.fingerprints import CircularFingerprint

@@ -49,16 +52,23 @@ class TestFeaturizedSamples(unittest.TestCase):
    samples = featurizer.featurize(input_file, self.feature_dir, self.samples_dir)

    # Splits featurized samples into train/test
    assert splittype in ["random", "specified", "scaffold"]
    if splittype == "random":
      splitter = RandomSplitter()
    elif splittype == "specified":
      splitter = SpecifiedSplitter()
    elif splittype == "scaffold":
      splitter = ScaffoldSplitter()
    if frac_valid > 0:
      train_samples, valid_samples, test_samples = samples.train_valid_test_split(
          splittype, train_dir=self.train_dir, valid_dir=self.valid_dir,
      train_samples, valid_samples, test_samples = splitter.train_valid_test_split(
          samples, train_dir=self.train_dir, valid_dir=self.valid_dir,
          test_dir=self.test_dir, frac_train=frac_train,
          frac_valid=frac_valid, frac_test=frac_test)

      return train_samples, valid_samples, test_samples
    else:
      train_samples, test_samples = samples.train_test_split(
          splittype, train_dir=self.train_dir, test_dir=self.test_dir,
      train_samples, test_samples = splitter.train_test_split(
          samples, train_dir=self.train_dir, test_dir=self.test_dir,
          frac_train=frac_train)
      return train_samples, test_samples

Loading