Commit c17663cb authored by Bharath Ramsundar's avatar Bharath Ramsundar
Browse files

Preliminary commit of task splits.

parent 28e3268c
Loading
Loading
Loading
Loading
+15 −15
Original line number Diff line number Diff line
@@ -237,7 +237,7 @@ class NumpyDataset(Dataset):

  def get_task_names(self):
    """Get the names of the tasks associated with this dataset."""
    tasks = np.arange(self._y.shape[1])
    return np.arange(self._y.shape[1])

  @property
  def X(self):
@@ -486,7 +486,6 @@ class DiskDataset(Dataset):
    """
    return self.metadata_df.shape[0]


  def itershards(self):
    """
    Return an object that iterates over all shards in dataset.
@@ -603,6 +602,7 @@ class DiskDataset(Dataset):
  @staticmethod
  def from_numpy(data_dir, X, y, w=None, ids=None, tasks=None, verbosity=None,
                 compute_feature_statistics=True):
    """Creates a DiskDataset object from specified Numpy arrays."""
    n_samples = len(X)
    # The -1 indicates that y will be reshaped to have length -1
    if n_samples > 0:
+2 −3
Original line number Diff line number Diff line
@@ -68,8 +68,8 @@ class Splitter(object):
      fold_datasets.append(fold_dataset)
    return fold_datasets

  def train_valid_test_split(self, dataset, train_dir,
                             valid_dir, test_dir, frac_train=.8,
  def train_valid_test_split(self, dataset, train_dir=None,
                             valid_dir=None, test_dir=None, frac_train=.8,
                             frac_valid=.1, frac_test=.1, seed=None,
                             log_every_n=1000,
                             compute_feature_statistics=True):
@@ -243,7 +243,6 @@ class RandomStratifiedSplitter(Splitter):
    return fold_datasets



class MolecularWeightSplitter(Splitter):
  """
  Class for doing data splits by molecular weight.
+77 −0
Original line number Diff line number Diff line
"""
Contains an abstract base class that supports chemically aware data splits.
"""
from __future__ import print_function
from __future__ import division
from __future__ import unicode_literals

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

import tempfile
import numpy as np
from rdkit import Chem
from deepchem.utils import ScaffoldGenerator
from deepchem.utils.save import log
from deepchem.datasets import NumpyDataset
from deepchem.featurizers.featurize import load_data
from deepchem.splits import Splitter

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.
  """

  def __init__(self):
    "Creates Task Splitter object."
    pass

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

    Parameters
    ----------
    dataset: deepchem.datasets.Dataset
      Dataset to be split
    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.
    """
    n_tasks = len(dataset.get_task_names())
    n_train = np.round(frac_train * n_tasks)
    n_valid = np.round(frac_valid * n_tasks)
    n_test = np.round(frac_test * n_tasks)
    if n_train + n_valid + n_test != n_tasks:
      raise ValueError("Train/Valid/Test fractions don't split tasks evenly.")
    ########################################### DEBUG
    print("train_valid_test_split")
    print("n_train, n_valid, n_test")
    print(n_train, n_valid, n_test)
    ########################################### DEBUG

    X, y, w, ids = dataset.X, dataset.y, dataset.w, dataset.ids
    
    train_dataset = NumpyDataset(X, y[:,:n_train], w[:,:n_train], ids)
    valid_dataset = NumpyDataset(
        X, y[:,n_train:n_train+n_valid], w[:,n_train:n_train+n_valid], ids)
    test_dataset = NumpyDataset(
        X, y[:,n_train+n_valid:], w[:,n_train+n_valid:], ids)
    ########################################### DEBUG
    print("train_dataset.get_task_names()")
    print(train_dataset.get_task_names())
    print("valid_dataset.get_task_names()")
    print(valid_dataset.get_task_names())
    print("test_dataset.get_task_names()")
    print(test_dataset.get_task_names())
    ########################################### DEBUG
    return train_dataset, valid_dataset, test_dataset
+47 −0
Original line number Diff line number Diff line

"""
Tests for splitter objects.
"""
from __future__ import division
from __future__ import print_function
from __future__ import unicode_literals

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

import tempfile
import numpy as np
from deepchem.splits.task_splitter import TaskSplitter
from deepchem.datasets import NumpyDataset
from deepchem.datasets.tests import TestDatasetAPI


class TestTaskSplitters(TestDatasetAPI):
  """
  Test some basic splitters.
  """

  def test_multitask_train_valid_test_split(self):
    """
    Test TaskSplitter train/valid/test split on multitask dataset.
    """
    n_samples = 100
    n_features = 10
    n_tasks = 10
    X = np.random.rand(n_samples, n_features)
    p = .05 # proportion actives
    y = np.random.binomial(1, p, size=(n_samples, n_tasks))
    dataset = NumpyDataset(X, y)
    ########################################### DEBUG
    print("dataset")
    print(dataset)
    ########################################### DEBUG

    task_splitter = TaskSplitter()
    train, valid, test = task_splitter.train_valid_test_split(
        dataset, frac_train=.4, frac_valid=.3, frac_test=.3)

    assert len(train.get_task_names()) == 4
    assert len(valid.get_task_names()) == 3
    assert len(test.get_task_names()) == 3