Commit 0017a775 authored by Aneesh Pappu's avatar Aneesh Pappu
Browse files

fixing and adding tests

parent 934c1672
Loading
Loading
Loading
Loading
+3 −3
Original line number Diff line number Diff line
@@ -5,7 +5,7 @@ from __future__ import print_function
from __future__ import division
from __future__ import unicode_literals

__author__ = "Bharath Ramsundar"
__author__ = "Bharath Ramsundar, Aneesh Pappu"
__copyright__ = "Copyright 2016, Stanford University"
__license__ = "GPL"

@@ -85,7 +85,7 @@ class StratifiedSplitter(Splitter):
  Class for doing stratified splits -- where data is too sparse to do regular splits
  """

  def __randomizeArrays(self, array_list):
  def __randomize_arrays(self, array_list):
    generator_state = np.random.get_state()
    for array in array_list:
      np.random.shuffle(array)
@@ -128,7 +128,7 @@ class StratifiedSplitter(Splitter):
   # Obtain original x, y, and w arrays
    numpyArrayList = dataset.to_numpy();

    numpyArrayList = self.__randomizeArrays(numpyArrayList)
    numpyArrayList = self.__randomize_arrays(numpyArrayList)
    X = numpyArrayList[0]
    y = numpyArrayList[1]
    w = numpyArrayList[2]
+13 −0
Original line number Diff line number Diff line
@@ -15,6 +15,8 @@ from deepchem.splits import RandomSplitter
from deepchem.splits import ScaffoldSplitter
from deepchem.splits import StratifiedSplitter
from deepchem.datasets.tests import TestDatasetAPI
from deepchem.datasets import Dataset
import pandas as pd

class TestSplitters(TestDatasetAPI):
  """
@@ -90,5 +92,16 @@ class TestSplitters(TestDatasetAPI):
          self.train_dir, self.valid_dir, self.test_dir,
          frac_train = 0.8, frac_valid = 0.1, frac_test = 0.1
      )
   train_np_list = train_data.to_numpy()
   y = train_data[1]
   #verify that each task in the train dataset has some hits
   y_df = pd.DataFrame(data = y)
   totalRows = len(y_df.index)
   for col in y_df:
       column = y_df[col]
       NaN_count = column.isnull().sum()
       numRows = len(colum)
       if NaN_count == totalRows:
           assert NaN_count != totalRows
   print("end of stratified test")
   assert 1 == 1