Commit 52761da8 authored by Aneesh Pappu's avatar Aneesh Pappu
Browse files

sparse stuff

parent 3e230064
Loading
Loading
Loading
Loading
+16 −0
Original line number Diff line number Diff line
@@ -73,3 +73,19 @@ class TestDatasetAPI(TestAPI):
        featurizer=featurizer,
        verbosity="low")
    return loader.featurize(input_file, self.data_dir)

  def load_sparse_multitask_dataset(self):
    """Load sparse tox multitask data."""
    if os.path.exists(self.data_dir):
      shutil.rmtree(self.data_dir)
    featurizer = CircularFingerprint(size=1024)
    tasks = ["task1", "task2", "task3", "task4", "task5", "task6",
             "task7", "task8", "task9"]
    input_file = os.path.join(
        self.current_dir, "../../models/tests/sparse_multitask_example.csv")
    loader = DataLoader(
        tasks=tasks,
        smiles_field="smiles",
        featurizers=featurizer,
        verbosity="low")
    return loader.featurize(input_file, self.data_dir)
+601 −0

File added.

Preview size limit exceeded, changes collapsed.

+15 −0
Original line number Diff line number Diff line
@@ -76,6 +76,21 @@ class Splitter(object):
    """
    raise NotImplementedError

class StratifiedSplitter(Splitter):
  """
  Class for doing stratified splits -- where data is too sparse to do regular splits
  """
  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):
   """
   Overrides parent implementation to do stratified split
   """
   numpyArrayList = dataset.to_numpy();
   print(numpyArrayList)
   return (False, False, False) #placeholder until rest of code is written

class MolecularWeightSplitter(Splitter):
  """
  Class for doing data splits by molecular weight.
+18 −4
Original line number Diff line number Diff line
@@ -19,8 +19,9 @@ class TestSplitters(TestDatasetAPI):
  """
  Test some basic splitters.
  """
  """
  def test_singletask_random_split(self):
    """Test singletask RandomSplitter class."""
    Test singletask RandomSplitter class.
    solubility_dataset = self.load_solubility_data()
    random_splitter = RandomSplitter()
    train_data, valid_data, test_data = \
@@ -33,7 +34,7 @@ class TestSplitters(TestDatasetAPI):
    assert len(test_data) == 1

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

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

  def test_multitask_scaffold_split(self):
    """Test multitask ScaffoldSplitter class."""
    Test multitask ScaffoldSplitter class.
    multitask_dataset = self.load_multitask_data()
    scaffold_splitter = ScaffoldSplitter()
    train_data, valid_data, test_data = \
@@ -70,3 +71,16 @@ 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()
   stratified_splitter = StratifiedSplitter()
   train_data, valid_data, test_data = \
      stratified_splitter.train_valid_test_split(
          sparse_dataset,
          self.train_dir, self.valid_dir, self.test_dir,
          frac_train = 0.8, frac_valid = 0.1, frac_test = 0.1
      )
   print("end of stratified test")
   assert 1 == 1