Commit 5b8679fd authored by Aneesh Pappu's avatar Aneesh Pappu
Browse files

stratified splitter code

parent c88cf853
Loading
Loading
Loading
Loading
+1 −1
Original line number Diff line number Diff line
@@ -75,7 +75,7 @@ class TestDatasetAPI(TestAPI):
    return loader.featurize(input_file, self.data_dir)

  def load_sparse_multitask_dataset(self):
    """Load sparse tox multitask data."""
    """Load sparse tox multitask data, sample dataset."""
    if os.path.exists(self.data_dir):
      shutil.rmtree(self.data_dir)
    featurizer = CircularFingerprint(size=1024)
+129 −4
Original line number Diff line number Diff line
@@ -11,12 +11,14 @@ __license__ = "GPL"

import os
import numpy as np
import pandas as pd
from rdkit import Chem
from deepchem.utils import ScaffoldGenerator
from deepchem.utils.save import log
from deepchem.datasets import Dataset
from deepchem.featurizers.featurize import load_data


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


class Splitter(object):
  """
  Abstract base class for chemically aware splits..
@@ -76,20 +79,142 @@ class Splitter(object):
    """
    raise NotImplementedError


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

  def __randomizeArrays(self, arraylist):
    generator_state = numpy.random.get_state()
    for array in arrayList:
      numpy.random.shuffle(array)
      numpy.random.set_state(generator_state)
    return arrayList

  def __generate_required_hits(self, y_df, frac_train):
    colIndex = 0
    required_hit_dict = {}
    totalCount = len(y_df.index)
    for col in y_df:
      NaN_count = y_df[col].isnull().sum()
      notNaN = totalCount - NaN_count
      requiredNotNaN = frac_train * notNaN
      required_hit_dict[colIndex] = requiredNotNaN
      colIndex += 1
    return required_hit_dict

  def __generate_required_index(self, y_df, required_hit_dict):
    index_dict = {}
    colIndex = 0
    for col in y_df:
      column = y_df[col]
      num_hit = 0
      num_required = required_hit_dict[colIndex]
      colIndex += 1
      for index, value in y_df[col].iteritems():
        if pd.notnull(value):
          num_hit += 1
          # check to see if number of hits has been hit
          if num_hit >= num_required:
            index_dict[colIndex] = index
            break
    return index_dict

  def train_valid_test_split(self, dataset, train_dir,
                             valid_dir, test_dir, frac_train=.8,
                             frac_valid=.1, frac_test=.1, seed=None,
                             log_every_n=1000):
   # Obtain original x, y, and w arrays
    numpyArrayList = dataset.to_numpy();

    numpyArrayList = randomizeArrays(numpyArrayList)
    X = numpyArrayList[0]
    y = numpyArrayList[1]
    w = numpyArrayList[2]
    ids = numpyArrayList[3]

    """
   Overrides parent implementation to do stratified split
    frac_train identifies percentage of datapoints that need to be present in split -- so 80% training data may actually be 90% of data (but 80% of actual datapoints, not NaN, will be present in split)
    """
   numpyArrayList = dataset.to_numpy();
   print(numpyArrayList)
   return (False, False, False) #placeholder until rest of code is written
    # find, for each task, the total number of hits and calculate the required
    # number of hits for valid split based on frac_train
    x_df = pd.DataFrame(data=x)
    y_df = pd.DataFrame(data=y)
    w_df = pd.DataFrame(data=w)
    id_df = pd.DataFrame(data=ids)

    required_hit_dict = __generate_required_hits(y_df, frac_train)
    index_dict = __generate_required_index(y_df, required_hit_dict)
    X_train, X_test, y_train, y_test, w_train, w_test, id_train, id_test = []

    # cycle through rows in y, copy over rows as appropriate
    for rowIndex, row in y_df.iterrows():
     weight_row = w_df.iloc[rowIndex].tolist() #get corresponding weight row as list
     weight_train_row = []
     weight_test_row = []
     for index, value in row.iteritems():
       # test if should be test or train data
       if rowIndex <= index_dict[index]: #train data
         weight_train_row.append(weight_row[index]) #add corresponding weight
         weight_test_row.append(0)
       else: #index is past test index -- this datapoint is test data
         weight_train_row.append(0)
         weight_test_row.append(weight_row[index])
     x_row = x_df.iloc[rowIndex].tolist()
     id_row = id_df.iloc[rowIndex].tolist()
     # check to see if any weight vectors are just zero
     if weight_train_row.count(0) == len(weight_train_row): #entire example is a test example
       # Add entire row to appropriate test arrays
       X_test.append(x_row) #get corresponding row from original x df
       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
       X_train.append(x_row)
       y_train.append(row)
       w_train.append(weight_train_row)
       id_train.append(id_row)
     else: #hybrid example -- feature X, results y, and smiles id are appended to both test and train. Weight vectors for train and row are appended as appropriately to dictate whether value is train or test
       X_train.append(x_row)
       X_test.append(x_row)
       y_train.append(row)
       y_test.append(row)
       w_train.append(weight_train_row)
       w_test.append(weight_test_row)
       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)
    y_test_np = np.array(y_test)
    w_train_np = np.array(w_train)
    w_test_np = np.array(w_test)
    id_train_np = np.array(id_train)
    id_test_np = np.array(id_test)

    # make valid split - 50/50 split of test
    X_split_list = np.vsplit(X_test_np, 2)
    y_split_list = np.vsplit(y_test_np, 2)
    w_split_list = np.vsplit(w_test_np, 2)
    id_split_list = np.vsplit(id_test_np, 2)

    X_test_np = X_split_list[0]
    X_valid_np = X_split_list[1]
    y_test_np = y_split_list[0]
    y_valid_np = y_split_list[1]
    w_test_np = w_split_list[0]
    w_valid_np = w_split_list[1]
    id_test_np = id_split_list[0]
    id_valid_np = id_split_list[1]

    # turn back into dataset objects
    train_data = Dataset.from_numpy(train_dir, X_train_np, y_train_np, w_train_np, id_train_np)
    valid_data = Dataset.from_numpy(valid_dir, X_valid_np, y_valid_np, w_valid_np, id_valid_np)
    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):
  """
+9 −2
Original line number Diff line number Diff line
@@ -20,9 +20,10 @@ class TestSplitters(TestDatasetAPI):
  """
  Test some basic splitters.
  """
  """
  def test_singletask_random_split(self):
    """
    Test singletask RandomSplitter class.
    """
    solubility_dataset = self.load_solubility_data()
    random_splitter = RandomSplitter()
    train_data, valid_data, test_data = \
@@ -35,7 +36,9 @@ class TestSplitters(TestDatasetAPI):
    assert len(test_data) == 1

  def test_singletask_scaffold_split(self):
    """
    Test singletask ScaffoldSplitter class.
    """
    solubility_dataset = self.load_solubility_data()
    scaffold_splitter = ScaffoldSplitter()
    train_data, valid_data, test_data = \
@@ -48,7 +51,9 @@ class TestSplitters(TestDatasetAPI):
    assert len(test_data) == 1

  def test_multitask_random_split(self):
    """
    Test multitask RandomSplitter class.
    """
    multitask_dataset = self.load_multitask_data()
    random_splitter = RandomSplitter()
    train_data, valid_data, test_data = \
@@ -61,7 +66,9 @@ class TestSplitters(TestDatasetAPI):
    assert len(test_data) == 1

  def test_multitask_scaffold_split(self):
    """
    Test multitask ScaffoldSplitter class.
    """
    multitask_dataset = self.load_multitask_data()
    scaffold_splitter = ScaffoldSplitter()
    train_data, valid_data, test_data = \
@@ -72,7 +79,7 @@ class TestSplitters(TestDatasetAPI):
    assert len(train_data) == 8
    assert len(valid_data) == 1
    assert len(test_data) == 1
"""

  def test_stratified_multitask_split(self):
   print("In stratified tester")
   sparse_dataset = self.load_sparse_multitask_dataset()