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

Some fixes to tests

parent a6305f1a
Loading
Loading
Loading
Loading
+0 −1
Original line number Diff line number Diff line
@@ -17,7 +17,6 @@ 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
+12 −2
Original line number Diff line number Diff line
@@ -28,6 +28,9 @@ from deepchem.transformers import LogTransformer
from deepchem.transformers import ClippingTransformer
from deepchem.hyperparameters import HyperparamOpt
from sklearn.ensemble import RandomForestRegressor
from deepchem.splits import RandomSplitter
from deepchem.splits import ScaffoldSplitter
from deepchem.splits import SpecifiedSplitter

class TestAPI(unittest.TestCase):
  """
@@ -114,8 +117,15 @@ class TestAPI(unittest.TestCase):
                                   shard_size=shard_size)

    # Splits featurized samples into train/test
    train_samples, test_samples = samples.train_test_split(
        splittype, self.train_dir, self.test_dir)
    assert splittype in ["random", "specified", "scaffold"]
    if splittype == "random":
      splitter = RandomSplitter()
    elif splittype == "specified":
      splitter = SpecifiedSplitter()
    elif splittype == "scaffold":
      splitter = ScaffoldSplitter()
    train_samples, test_samples = splitter.train_test_split(
        samples, self.train_dir, self.test_dir)

    use_user_specified_features = (user_specified_features is not None)
    train_dataset = Dataset(data_dir=self.train_dir, samples=train_samples, 
+7 −4
Original line number Diff line number Diff line
@@ -12,6 +12,8 @@ __license__ = "LGPL"
import os
import numpy as np
from rdkit import Chem
from deepchem.utils import ScaffoldGenerator
from deepchem.utils.save import log
from deepchem.featurizers.featurize import FeaturizedSamples

def generate_scaffold(smiles, include_chirality=False):
@@ -73,7 +75,7 @@ class Splitter(object):

    return train_samples, valid_samples, test_samples

  def train_test_split(self, samples, splittype, train_dir, test_dir, seed=None,
  def train_test_split(self, samples, train_dir, test_dir, seed=None,
                       frac_train=.8, reload=False):
    """
    Splits self into train/test sets.
@@ -81,7 +83,7 @@ class Splitter(object):
    Returns FeaturizedDataset objects.
    """
    train_samples, _, test_samples = self.train_valid_test_split(
        samples, splittype, train_dir, valid_dir=None, test_dir=test_dir,
        samples, 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
@@ -114,7 +116,8 @@ 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):
  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.
    """
@@ -147,7 +150,7 @@ class ScaffoldSplitter(Splitter):
        train_inds += scaffold_set
    return train_inds, valid_inds, test_inds

class SpecifiedSplit(Splitter):
class SpecifiedSplitter(Splitter):
  """
  Class that splits data according to user specification.
  """
+2 −2
Original line number Diff line number Diff line
@@ -91,7 +91,7 @@ class TestSplitAPI(unittest.TestCase):
        input_transformer_classes, output_transformer_classes,
        input_file, task_types.keys())

  def _load_multitask_data(self):
  def _load_multitask_samples(self):
    """Load example multitask data."""
    compound_featurizers = [CircularFingerprint(size=1024)]
    complex_featurizers = []
@@ -102,7 +102,7 @@ class TestSplitAPI(unittest.TestCase):
             "task13", "task14", "task15", "task16"]
    task_types = {task: "classification" for task in tasks}
    input_file = "multitask_example.csv"
    return self._create_dataset(
    return self._gen_samples(
        compound_featurizers, complex_featurizers,
        input_transformer_classes, output_transformer_classes,
        input_file, task_types.keys())
+41 −1
Original line number Diff line number Diff line
@@ -12,6 +12,7 @@ __license__ = "LGPL"
import os
import unittest
from deepchem.splits import RandomSplitter
from deepchem.splits import ScaffoldSplitter
from deepchem.splits.tests import TestSplitAPI

class TestSplitters(TestSplitAPI):
@@ -19,7 +20,7 @@ class TestSplitters(TestSplitAPI):
  Test some basic splitters.
  """
  def test_singletask_random_split(self):
    """Test RandomSplitter class."""
    """Test singletask RandomSplitter class."""
    solubility_samples = self._load_solubility_samples()
    random_splitter = RandomSplitter()
    train_data, valid_data, test_data = \
@@ -30,3 +31,42 @@ class TestSplitters(TestSplitAPI):
    assert len(train_data) == 8
    assert len(valid_data) == 1
    assert len(test_data) == 1

  def test_singletask_scaffold_split(self):
    """Test singletask ScaffoldSplitter class."""
    solubility_samples = self._load_solubility_samples()
    scaffold_splitter = ScaffoldSplitter()
    train_data, valid_data, test_data = \
        scaffold_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

  def test_multitask_random_split(self):  
    """Test multitask RandomSplitter class."""
    multitask_samples = self._load_multitask_samples()
    random_splitter = RandomSplitter()
    train_data, valid_data, test_data = \
        random_splitter.train_valid_test_split(
            multitask_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

  def test_multitask_scaffold_split(self):  
    """Test multitask ScaffoldSplitter class."""
    multitask_samples = self._load_multitask_samples()
    scaffold_splitter = ScaffoldSplitter()
    train_data, valid_data, test_data = \
        scaffold_splitter.train_valid_test_split(
            multitask_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