Commit dd3862da authored by Aneesh Pappu's avatar Aneesh Pappu
Browse files

fixing tests

parent 0017a775
Loading
Loading
Loading
Loading
+9 −3
Original line number Diff line number Diff line
@@ -171,7 +171,8 @@ class StratifiedSplitter(Splitter):
                y_test.append(row)
                w_test.append(weight_test_row)
                id_test.append(id_row)
     elif weight_test_row.count(0) == len(weight_test_row): #entirely train example
            elif weight_test_row.count(0) == len(weight_test_row):
                # entirely train example
                X_train.append(x_row)
                y_train.append(row)
                w_train.append(weight_train_row)
@@ -186,7 +187,6 @@ class StratifiedSplitter(Splitter):
                id_train.append(id_row)
                id_test.append(id_row)


        X_train_np = np.array(X_train)
        X_test_np = np.array(X_test)
        y_train_np = np.array(y_train)
@@ -217,10 +217,12 @@ class StratifiedSplitter(Splitter):
        test_data = Dataset.from_numpy(test_dir, X_test_np, y_test_np, w_test_np, id_test_np)
        return (train_data, valid_data, test_data)


class MolecularWeightSplitter(Splitter):
    """
  Class for doing data splits by molecular weight.
  """

    def split(self, dataset, seed=None, frac_train=.8, frac_valid=.1,
              frac_test=.1, log_every_n=None):
        """
@@ -247,15 +249,18 @@ class MolecularWeightSplitter(Splitter):
        return (sortidx[:train_cutoff], sortidx[train_cutoff:valid_cutoff],
                sortidx[valid_cutoff:])


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

    def split(self, dataset, seed=None, frac_train=.8, frac_valid=.1,
              frac_test=.1, log_every_n=None):
        """
    Splits internal compounds randomly into train/validation/test.
    """
<<<<<<< 0017a77551376a895531b1020b82acc9029563cd
    np.testing.assert_almost_equal(frac_train + frac_valid + frac_test, 1.)
    np.random.seed(seed)
    num_datapoints = len(dataset)
@@ -269,11 +274,13 @@ class ScaffoldSplitter(Splitter):
    """
  Class for doing data splits based on the scaffold of small molecules.
  """

    def split(self, dataset, frac_train=.8, frac_valid=.1, frac_test=.1,
              log_every_n=1000):
        """
    Splits internal compounds into train/validation/test by scaffold.
    """
<<<<<<< 0017a77551376a895531b1020b82acc9029563cd
    np.testing.assert_almost_equal(frac_train + frac_valid + frac_test, 1.)
    scaffolds = {}
    log("About to generate scaffolds", self.verbosity)
@@ -313,7 +320,6 @@ class SpecifiedSplitter(Splitter):
    raw_df = load_data([input_file], shard_size=None).next()
    self.splits = raw_df[split_field].values
    self.verbosity = verbosity

  def split(self, dataset, frac_train=.8, frac_valid=.1, frac_test=.1,
                  log_every_n=1000):
  """
+3 −1
Original line number Diff line number Diff line
@@ -18,10 +18,12 @@ from deepchem.datasets.tests import TestDatasetAPI
from deepchem.datasets import Dataset
import pandas as pd


class TestSplitters(TestDatasetAPI):
    """
  Test some basic splitters.
  """

    def test_singletask_random_split(self):
        """
    Test singletask RandomSplitter class.
@@ -93,7 +95,7 @@ class TestSplitters(TestDatasetAPI):
                frac_train=0.8, frac_valid=0.1, frac_test=0.1
            )
        train_np_list = train_data.to_numpy()
   y = train_data[1]
        y = train_np_list[1]
        # verify that each task in the train dataset has some hits
        y_df = pd.DataFrame(data=y)
        totalRows = len(y_df.index)