Commit 5aaee3ff authored by Bharath Ramsundar's avatar Bharath Ramsundar
Browse files

Adding initial (failing) multitask hyperparam opt test

parent 21555864
Loading
Loading
Loading
Loading
+36 −0
Original line number Diff line number Diff line
@@ -62,3 +62,39 @@ class TestHyperparamOptAPI(TestAPI):
    self._hyperparam_opt(rf_model_builder, params_dict, train_dataset,
                         valid_dataset, output_transformers, task_types,
                         metric)

  def test_multitask_keras_mlp_ECFP_classification_hyperparam_opt(self):
    """Straightforward test of Keras multitask deepchem classification API."""
    from deepchem.models.keras_models.fcnet import MultiTaskDNN
    splittype = "scaffold"
    output_transformers = []
    input_transformers = []
    task_type = "classification"
    # TODO(rbharath): There should be some automatic check to ensure that all
    # required model_params are specified.
    model_params = {"nb_hidden": 10, "activation": "relu",
                    "dropout": .5, "learning_rate": .01,
                    "momentum": .9, "nesterov": False,
                    "decay": 1e-4, "batch_size": 5,
                    "nb_epoch": 2, "init": "glorot_uniform",
                    "nb_layers": 1, "batchnorm": False}

    input_file = os.path.join(self.current_dir, "multitask_example.csv")
    tasks = ["task0", "task1", "task2", "task3", "task4", "task5", "task6",
             "task7", "task8", "task9", "task10", "task11", "task12",
             "task13", "task14", "task15", "task16"]
    task_types = {task: task_type for task in tasks}

    compound_featurizers = [CircularFingerprint(size=1024)]
    complex_featurizers = []

    train_dataset, test_dataset, _, transformers = self._featurize_train_test_split(
        splittype, compound_featurizers, 
        complex_featurizers, input_transformers,
        output_transformers, input_file, task_types.keys())
    model_params["data_shape"] = train_dataset.get_data_shape()
    metric = Metric(metrics.mean_roc_auc_score)
    
    self._hyperparam_opt(MultiTaskDNN, params_dict, train_dataset,
                         valid_dataset, output_transformers, task_types,
                         metric)
+0 −3
Original line number Diff line number Diff line
@@ -15,9 +15,6 @@
# limitations under the License.
"""Evaluation metrics."""

import collections


import numpy as np
import warnings
from sklearn.metrics import roc_auc_score