Commit 110dc734 authored by Bharath Ramsundar's avatar Bharath Ramsundar
Browse files

Changes

parent 3a70ec24
Loading
Loading
Loading
Loading
+16 −5
Original line number Diff line number Diff line
"""
Gathers all splitters in one place for convenient imports
"""
# TODO(rbharath): Get rid of * import
from deepchem.splits.splitters import *
from deepchem.splits.splitters import ScaffoldSplitter
from deepchem.splits.splitters import SpecifiedSplitter
from deepchem.splits.splitters import generate_scaffold
from deepchem.splits.splitters import randomize_arrays
from deepchem.splits.splitters import Splitter
from deepchem.splits.splitters import RandomGroupSplitter
from deepchem.splits.splitters import RandomStratifiedSplitter
from deepchem.splits.splitters import SingletaskStratifiedSplitter
from deepchem.splits.splitters import MolecularWeightSplitter
from deepchem.splits.splitters import MaxMinSplitter
from deepchem.splits.splitters import RandomSplitter
from deepchem.splits.splitters import IndexSplitter
from deepchem.splits.splitters import IndiceSplitter
from deepchem.splits.splitters import RandomGroupSplitter
from deepchem.splits.splitters import ClusterFps
from deepchem.splits.splitters import ButinaSplitter
from deepchem.splits.splitters import ScaffoldSplitter
from deepchem.splits.splitters import FingerprintSplitter
from deepchem.splits.splitters import SpecifiedSplitter
from deepchem.splits.splitters import FingerprintSplitter
from deepchem.splits.splitters import TimeSplitterPDBbind
from deepchem.splits.task_splitter import merge_fold_datasets
from deepchem.splits.task_splitter import TaskSplitter
+272 −92

File changed.

Preview size limit exceeded, changes collapsed.

+36 −16
Original line number Diff line number Diff line
@@ -8,15 +8,21 @@ from deepchem.utils import ScaffoldGenerator
from deepchem.data import NumpyDataset
from deepchem.utils.save import load_data
from deepchem.splits import Splitter
from deepchem.utils.data import datasetify

logger = logging.getLogger(__name__)

def merge_fold_datasets(fold_datasets):
  """Merges fold datasets together.

  Assumes that fold_datasets were outputted from k_fold_split. Specifically,
  assumes that each dataset contains the same datapoints, listed in the same
  ordering.
  Assumes that fold_datasets were outputted from k_fold_split.
  Specifically, assumes that each dataset contains the same
  datapoints, listed in the same ordering.

  Parameters
  ----------
  fold_dataset: list[dc.data.Dataset]
    Each entry of this list should be a `dc.data.Dataset` object.
  """
  if not len(fold_datasets):
    return None
@@ -38,36 +44,43 @@ class TaskSplitter(Splitter):
  """
  Provides a simple interface for splitting datasets task-wise.

  For some learning problems, the training and test datasets should
  have different tasks entirely. This is a different paradigm from the
  usual Splitter, which ensures that split datasets have different
  datapoints, not different tasks.
  For some learning problems, the training and test datasets
  should have different tasks entirely. This is a different
  paradigm from the usual Splitter, which ensures that split
  datasets have different datapoints, not different tasks.
  """

  def __init__(self):
    "Creates Task Splitter object."
    pass
  def __init__(self, *args, **kwargs):
    """Creates Task Splitter object."""
    super(TaskSplitter, self).__init__(*args, **kwargs)

  def train_valid_test_split(self,
                             dataset,
                             frac_train=.8,
                             frac_valid=.1,
                             frac_test=.1):
                             frac_test=.1,
                             seed=None):
    """Performs a train/valid/test split of the tasks for dataset.

    If split is uneven, spillover goes to test.

    Parameters
    ----------
    dataset: dc.data.Dataset
      Dataset to be split
    dataset: data-like object. 
      Dataset to do a k-fold split on. This should either be of type
      `dc.data.Dataset` or a type that `dc.utils.data.datasetify` can
      convert into a `Dataset`.
    frac_train: float, optional
      Proportion of tasks to be put into train. Rounded to nearest int.
    frac_valid: float, optional
      Proportion of tasks to be put into valid. Rounded to nearest int.
    frac_test: float, optional
      Proportion of tasks to be put into test. Rounded to nearest int.
    seed: int, optional
      Random seed to make the split deterministic 
    """
    np.random.seed(seed)
    dataset = datasetify(dataset)
    np.testing.assert_almost_equal(frac_train + frac_valid + frac_test, 1)
    n_tasks = len(dataset.get_task_names())
    n_train = int(np.round(frac_train * n_tasks))
@@ -83,18 +96,25 @@ class TaskSplitter(Splitter):
                                w[:, n_train + n_valid:], ids)
    return train_dataset, valid_dataset, test_dataset

  def k_fold_split(self, dataset, K):
  def k_fold_split(self, dataset, K, seed=None):
    """Performs a K-fold split of the tasks for dataset.

    If split is uneven, spillover goes to last fold.

    Parameters
    ----------
    dataset: dc.data.Dataset
      Dataset to be split
    dataset: data like object. 
      Dataset to be split. This should either be of type
      `dc.data.Dataset` or a type that `dc.utils.data.datasetify` can
      convert into a `Dataset`.
    K: int
      Number of splits to be made
    seed: int, optional
      Random seed to make the split deterministic 
    """
    if seed is not None:
      np.random.seed(seed)
    dataset = datasetify(dataset)
    n_tasks = len(dataset.get_task_names())
    n_per_fold = int(np.round(n_tasks / float(K)))
    if K * n_per_fold != n_tasks:
Loading