Commit 588801cf authored by leswing's avatar leswing
Browse files

Passing tests

parent 899f1ef9
Loading
Loading
Loading
Loading
+7 −0
Original line number Diff line number Diff line
@@ -395,6 +395,13 @@ class NumpyDataset(Dataset):
    ids = self.ids[indices]
    return NumpyDataset(X, y, w, ids)

  @staticmethod
  def from_DiskDataset(ds):
    """
    :param ds: DiskDataset
    :return: NumpyDataset with the same data
    """
    return NumpyDataset(ds.X, ds.y, ds.w, ds.ids)

class DiskDataset(Dataset):
  """
+11 −4
Original line number Diff line number Diff line
@@ -60,7 +60,10 @@ class Splitter(object):
      directories = [tempfile.mkdtemp() for _ in range(k)]
    else:
      assert len(directories) == k
    fold_datasets = []
    all_ids = dataset.ids
    cv_datasets = []
    train_ds_base = None
    train_datasets = []
    # rem_dataset is remaining portion of dataset
    rem_dataset = dataset
    for fold in range(k):
@@ -73,11 +76,15 @@ class Splitter(object):
        frac_train=frac_fold,
        frac_valid=1 - frac_fold,
        frac_test=0)
      fold_dataset = rem_dataset.select(fold_inds, fold_dir)
      cv_dataset = rem_dataset.select(fold_inds, fold_dir)
      rem_dir = tempfile.mkdtemp()
      rem_dataset = rem_dataset.select(rem_inds, rem_dir)
      fold_datasets.append(fold_dataset)
    return fold_datasets

      train_dataset = DiskDataset.merge(filter(lambda x: x is not None, [train_ds_base, cv_dataset]))
      train_datasets.append(train_dataset)
      train_ds_base = DiskDataset.merge(filter(lambda x: x is not None, [train_ds_base, cv_dataset]))
      cv_datasets.append(cv_dataset)
    return list(zip(train_datasets, cv_datasets))

  def train_valid_test_split(self,
                             dataset,
+11 −9
Original line number Diff line number Diff line
@@ -7,6 +7,7 @@ from __future__ import unicode_literals

from rdkit.Chem.Fingerprints import FingerprintMols


__author__ = "Bharath Ramsundar, Aneesh Pappu"
__copyright__ = "Copyright 2016, Stanford University"
__license__ = "MIT"
@@ -15,6 +16,7 @@ import tempfile
import unittest
import numpy as np
import deepchem as dc
from deepchem.data import NumpyDataset
from rdkit import Chem, DataStructs


@@ -133,7 +135,7 @@ class TestSplitters(unittest.TestCase):
    K = 5
    fold_datasets = random_splitter.k_fold_split(solubility_dataset, K)
    for fold in range(K):
      fold_dataset = fold_datasets[fold]
      fold_dataset = fold_datasets[fold][1]
      # Verify lengths is 10/k == 2
      assert len(fold_dataset) == 2
      # Verify that compounds in this fold are subset of original compounds
@@ -143,11 +145,11 @@ class TestSplitters(unittest.TestCase):
      for other_fold in range(K):
        if fold == other_fold:
          continue
        other_fold_dataset = fold_datasets[other_fold]
        other_fold_dataset = fold_datasets[other_fold][1]
        other_fold_ids_set = set(other_fold_dataset.ids)
        assert fold_ids_set.isdisjoint(other_fold_ids_set)

    merged_dataset = dc.data.DiskDataset.merge(fold_datasets)
    merged_dataset = dc.data.DiskDataset.merge([x[1] for x in fold_datasets])
    assert len(merged_dataset) == len(solubility_dataset)
    assert sorted(merged_dataset.ids) == (sorted(solubility_dataset.ids))

@@ -163,7 +165,7 @@ class TestSplitters(unittest.TestCase):
    fold_datasets = index_splitter.k_fold_split(solubility_dataset, K)

    for fold in range(K):
      fold_dataset = fold_datasets[fold]
      fold_dataset = fold_datasets[fold][1]
      # Verify lengths is 10/k == 2
      assert len(fold_dataset) == 2
      # Verify that compounds in this fold are subset of original compounds
@@ -173,11 +175,11 @@ class TestSplitters(unittest.TestCase):
      for other_fold in range(K):
        if fold == other_fold:
          continue
        other_fold_dataset = fold_datasets[other_fold]
        other_fold_dataset = fold_datasets[other_fold][1]
        other_fold_ids_set = set(other_fold_dataset.ids)
        assert fold_ids_set.isdisjoint(other_fold_ids_set)

    merged_dataset = dc.data.DiskDataset.merge(fold_datasets)
    merged_dataset = dc.data.DiskDataset.merge([x[1] for x in fold_datasets])
    assert len(merged_dataset) == len(solubility_dataset)
    assert sorted(merged_dataset.ids) == (sorted(solubility_dataset.ids))

@@ -193,7 +195,7 @@ class TestSplitters(unittest.TestCase):
    fold_datasets = scaffold_splitter.k_fold_split(solubility_dataset, K)

    for fold in range(K):
      fold_dataset = fold_datasets[fold]
      fold_dataset = fold_datasets[fold][1]
      # Verify lengths is 10/k == 2
      assert len(fold_dataset) == 2
      # Verify that compounds in this fold are subset of original compounds
@@ -203,11 +205,11 @@ class TestSplitters(unittest.TestCase):
      for other_fold in range(K):
        if fold == other_fold:
          continue
        other_fold_dataset = fold_datasets[other_fold]
        other_fold_dataset = fold_datasets[other_fold][1]
        other_fold_ids_set = set(other_fold_dataset.ids)
        assert fold_ids_set.isdisjoint(other_fold_ids_set)

    merged_dataset = dc.data.DiskDataset.merge(fold_datasets)
    merged_dataset = dc.data.DiskDataset.merge([x[1] for x in fold_datasets])
    assert len(merged_dataset) == len(solubility_dataset)
    assert sorted(merged_dataset.ids) == (sorted(solubility_dataset.ids))