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

Forgot to add in new files

parent 1cb558d1
Loading
Loading
Loading
Loading
+169 −0
Original line number Diff line number Diff line
"""
Contains an abstract base class that supports chemically aware data splits.
"""
from __future__ import print_function
from __future__ import division
from __future__ import unicode_literals

__author__ = "Bharath Ramsundar"
__copyright__ = "Copyright 2016, Stanford University"
__license__ = "LGPL"

import os
import numpy as np
from rdkit import Chem
from deepchem.featurizers.featurize import FeaturizedSamples

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

class Splitter(object):
  """
  Abstract base class for chemically aware splits..
  """

  def __init__(self, verbosity=None):
    """Creates splitter object."""
    self.verbosity = verbosity

  def train_valid_test_split(self, samples, 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:
      train_inds, valid_inds, test_inds = self.split(
          samples,
          frac_train=frac_train, frac_test=frac_test,
          frac_valid=frac_valid, log_every_n=log_every_n)
    train_samples, valid_samples, test_samples = None, None, None
    dataset_files = samples.dataset_files
    if train_dir is not None:
      train_samples = FeaturizedSamples(samples_dir=train_dir, 
                                        dataset_files=dataset_files,
                                        featurizers=samples.featurizers,
                                        verbosity=self.verbosity,
                                        reload=False)
      if not reload:
        train_samples._set_compound_df(samples.compounds_df.iloc[train_inds])
    if test_dir is not None:
      test_samples = FeaturizedSamples(samples_dir=test_dir, 
                                       dataset_files=dataset_files,
                                       featurizers=samples.featurizers,
                                       verbosity=self.verbosity,
                                       reload=False)
      if not reload:
        test_samples._set_compound_df(samples.compounds_df.iloc[test_inds])
    if valid_dir is not None:
      valid_samples = FeaturizedSamples(samples_dir=valid_dir, 
                                       dataset_files=dataset_files,
                                       featurizers=samples.featurizers,
                                       verbosity=self.verbosity,
                                       reload=False)
      if not reload:
        valid_samples._set_compound_df(samples.compounds_df.iloc[valid_inds])

    return train_samples, valid_samples, test_samples

  def train_test_split(self, samples, 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(
        samples, 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 split(self, samples, frac_train=None, frac_valid=None, frac_test=None,
            log_every_n=None):
    """
    Stub to be filled in by child classes.
    """
    raise NotImplementedError

class RandomSplitter(Splitter):
  """
  Class for doing random data splits.
  """
  def split(self, samples, seed=None, frac_train=.8, frac_valid=.1,
            frac_test=.1, log_every_n=None):
    """
    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(samples.compounds_df)
    valid_cutoff = (frac_train+frac_valid) * len(samples.compounds_df)
    shuffled = np.random.permutation(range(len(samples.compounds_df)))
    return (shuffled[:train_cutoff], shuffled[train_cutoff:valid_cutoff],
            shuffled[valid_cutoff:])

class ScaffoldSplitter(Splitter):
  """
  Class for doing data splits based on the scaffold of small molecules.
  """
  def split(self, samples, 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 samples.compounds_df.iterrows():
      if self.verbosity is not None and ind % log_every_n == 0:
        log("Generating scaffold %d/%d" % (ind, len(samples.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(samples.compounds_df)
    valid_cutoff = (frac_train+frac_valid) * len(samples.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

class SpecifiedSplit(Splitter):
  """
  Class that splits data according to user specification.
  """
  def split(self, samples, frac_train=.8, frac_valid=.1, frac_test=.1,
            log_every_n=1000):
    """
    Splits internal compounds into train/validation/test by user-specification.
    """
    train_inds, valid_inds, test_inds = [], [], []
    for ind, row in samples.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
+109 −0
Original line number Diff line number Diff line
"""
General API for testing splitter objects
"""
from __future__ import print_function
from __future__ import division
from __future__ import unicode_literals

__author__ = "Bharath Ramsundar"
__copyright__ = "Copyright 2016, Stanford University"
__license__ = "LGPL"

import os
import shutil
import tempfile
import unittest
from deepchem.featurizers.featurize import DataFeaturizer
from deepchem.featurizers.fingerprints import CircularFingerprint

class TestSplitAPI(unittest.TestCase):
  """
  Test top-level API for Splitter 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()
    self.train_dir = tempfile.mkdtemp()
    self.valid_dir = tempfile.mkdtemp()
    self.test_dir = tempfile.mkdtemp()

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

  def _gen_samples(self, compound_featurizers, complex_featurizers,
                   input_transformer_classes, output_transformer_classes,
                   input_file, tasks,
                   protein_pdb_field=None, ligand_pdb_field=None,
                   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,
                                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)
    return samples

  def _load_solubility_samples(self):
    """Loads solubility data from example.csv"""
    compound_featurizers = [CircularFingerprint(size=1024)]
    complex_featurizers = []
    input_transformer_classes = []
    output_transformer_classes = []
    task_types = {"log-solubility": "regression"}
    input_file = "example.csv"
    return self._gen_samples(
        compound_featurizers, complex_featurizers,
        input_transformer_classes, output_transformer_classes,
        input_file, task_types.keys())

  def _load_classification_samples(self):
    """Loads classification data from example.csv"""
    compound_featurizers = [CircularFingerprint(size=1024)]
    complex_featurizers = []
    input_transformer_classes = []
    output_transformer_classes = []
    task_types = {"outcome": "classification"}
    input_file = "example_classification.csv"
    return self._gen_samples(
        compound_featurizers, complex_featurizers,
        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)]
    complex_featurizers = []
    output_transformer_classes = []
    input_transformer_classes = []
    tasks = ["task0", "task1", "task2", "task3", "task4", "task5", "task6",
             "task7", "task8", "task9", "task10", "task11", "task12",
             "task13", "task14", "task15", "task16"]
    task_types = {task: "classification" for task in tasks}
    input_file = "multitask_example.csv"
    return self._create_dataset(
        compound_featurizers, complex_featurizers,
        input_transformer_classes, output_transformer_classes,
        input_file, task_types.keys())
+32 −0
Original line number Diff line number Diff line
"""
Tests for splitter objects. 
"""
from __future__ import print_function
from __future__ import division
from __future__ import unicode_literals

__author__ = "Bharath Ramsundar"
__copyright__ = "Copyright 2016, Stanford University"
__license__ = "LGPL"

import os
import unittest
from deepchem.splits import RandomSplitter
from deepchem.splits.tests import TestSplitAPI

class TestSplitters(TestSplitAPI):
  """
  Test some basic splitters.
  """
  def test_singletask_random_split(self):
    """Test RandomSplitter class."""
    solubility_samples = self._load_solubility_samples()
    random_splitter = RandomSplitter()
    train_data, valid_data, test_data = \
        random_splitter.train_valid_test_split(
            solubility_samples,
            self.train_dir, self.valid_dir, self.test_dir,
            frac_train=0.8, frac_valid=0.1, frac_test=0.1, reload=False)
    assert len(train_data) == 8
    assert len(valid_data) == 1
    assert len(test_data) == 1