Commit 29cab76f authored by Bharath Ramsundar's avatar Bharath Ramsundar
Browse files

Add tasks parameter to models to ensure tasks retain ordering.

parent 57858df9
Loading
Loading
Loading
Loading
+5 −7
Original line number Diff line number Diff line
@@ -125,7 +125,7 @@ class Dataset(object):
    """
    if not len(self.metadata_df):
      raise ValueError("No data in dataset.")
    return sorted(self.metadata_df.iterrows().next()[1]['task_names'])
    return self.metadata_df.iterrows().next()[1]['task_names']

  def get_data_shape(self):
    """
@@ -309,8 +309,7 @@ def write_dataset_single(val, data_dir, feature_types, tasks):
  # TODO(rbharath): This is a hack. clean up.
  if not len(df):
    return None
  sorted_tasks = sorted(tasks)
  ids, X, y, w = _df_to_numpy(df, feature_types, sorted_tasks)
  ids, X, y, w = _df_to_numpy(df, feature_types, tasks)
  X_sums, X_sum_squares, X_n = compute_sums_and_nb_sample(X)
  y_sums, y_sum_squares, y_n = compute_sums_and_nb_sample(y, w)

@@ -346,7 +345,7 @@ def write_dataset_single(val, data_dir, feature_types, tasks):
  save_to_disk(ids, out_ids)
  # TODO(rbharath): Should X be saved to out_X_transformed as well? Since
  # itershards expects to loop over X-transformed? (Ditto for y/w)
  return([df_file, sorted_tasks, out_ids, out_X, out_X_transformed, out_y,
  return([df_file, tasks, out_ids, out_X, out_X_transformed, out_y,
          out_y_transformed, out_w, out_w_transformed,
          out_X_sums, out_X_sum_squares, out_X_n,
          out_y_sums, out_y_sum_squares, out_y_n])
@@ -360,10 +359,9 @@ def _df_to_numpy(df, feature_types, tasks):
        "Featurized data does not support requested feature_types.")
  # perform common train/test split across all tasks
  n_samples = df.shape[0]
  sorted_tasks = sorted(tasks)
  n_tasks = len(sorted_tasks)
  n_tasks = len(tasks)
  n_features = None
  y = df[sorted_tasks].values
  y = df[tasks].values
  y = np.reshape(y, (n_samples, n_tasks))
  w = np.ones((n_samples, n_tasks))
  missing = np.zeros_like(y).astype(int)
+9 −5
Original line number Diff line number Diff line
@@ -55,12 +55,14 @@ class TestDatasetAPI(TestSplitAPI):
    complex_featurizers = []
    input_transformer_classes = []
    output_transformer_classes = []
    task_types = {"log-solubility": "regression"}
    tasks = ["log-solubility"]
    task_type = "regression"
    task_types = {task: task_type for task in tasks}
    input_file = "example.csv"
    return self._create_dataset(
        compound_featurizers, complex_featurizers,
        input_transformer_classes, output_transformer_classes,
        input_file, task_types.keys())
        input_file, tasks)

  def _load_classification_data(self):
    """Loads classification data from example.csv"""
@@ -68,12 +70,14 @@ class TestDatasetAPI(TestSplitAPI):
    complex_featurizers = []
    input_transformer_classes = []
    output_transformer_classes = []
    task_types = {"outcome": "classification"}
    tasks = ["outcome"]
    task_type = "classification"
    task_types = {task: task_type for task in tasks}
    input_file = "example_classification.csv"
    return self._create_dataset(
        compound_featurizers, complex_featurizers,
        input_transformer_classes, output_transformer_classes,
        input_file, task_types.keys())
        input_file, tasks)

  def _load_multitask_data(self):
    """Load example multitask data."""
@@ -89,5 +93,5 @@ class TestDatasetAPI(TestSplitAPI):
    return self._create_dataset(
        compound_featurizers, complex_featurizers,
        input_transformer_classes, output_transformer_classes,
        input_file, task_types.keys())
        input_file, tasks)
+3 −3
Original line number Diff line number Diff line
@@ -91,7 +91,7 @@ class DataFeaturizer(object):
      raise ValueError("tasks must be a list.")
    assert verbosity in [None, "low", "high"]
    self.verbosity = verbosity
    self.sorted_tasks = sorted(tasks)
    self.tasks = tasks
    self.smiles_field = smiles_field
    self.split_field = split_field
    if id_field is None:
@@ -189,7 +189,7 @@ class DataFeaturizer(object):
    else:
      raise ValueError("Unrecognized input_type")
    if self.threshold is not None:
      for task in self.sorted_tasks:
      for task in self.tasks:
        raw = _process_field(data[task])
        if not isinstance(raw, float):
          raise ValueError("Cannot threshold non-float fields.")
@@ -201,7 +201,7 @@ class DataFeaturizer(object):
    df = pd.DataFrame(ori_df[[self.id_field]])
    df.columns = ["mol_id"]
    df["smiles"] = ori_df[[self.smiles_field]]
    for task in self.sorted_tasks:
    for task in self.tasks:
      df[task] = ori_df[[task]]
    if self.user_specified_features is not None:
      for feature in self.user_specified_features:
+16 −8
Original line number Diff line number Diff line
@@ -78,11 +78,13 @@ class TestFeaturizedSamples(unittest.TestCase):
    input_transforms = []
    output_transforms = ["normalize"]
    model_params = {}
    task_types = {"log-solubility": "regression"}
    tasks = ["log-solubility"]
    task_type = "regression"
    task_types = {task: task_type for task in tasks}
    input_file = "../../models/test/example.csv"
    train_samples, valid_samples, test_samples = (
        self._featurize_train_valid_test_split(
            splittype, input_file, task_types.keys(), frac_train=.8,
            splittype, input_file, tasks, frac_train=.8,
            frac_valid=.1, frac_test=.1))
    assert len(train_samples) == 8
    assert len(valid_samples) == 1
@@ -94,11 +96,13 @@ class TestFeaturizedSamples(unittest.TestCase):
    input_transforms = []
    output_transforms = ["normalize"]
    model_params = {}
    task_types = {"log-solubility": "regression"}
    tasks = ["log-solubility"]
    task_type = "regression"
    task_types = {task: task_type for task in tasks}
    input_file = "../../models/test/example.csv"
    train_samples, test_samples = (
        self._featurize_train_valid_test_split(
            splittype, input_file, task_types.keys(), frac_train=.8,
            splittype, input_file, tasks), frac_train=.8,
            frac_valid=0, frac_test=.2))
    assert len(train_samples) == 8
    assert len(test_samples) == 2
@@ -109,11 +113,13 @@ class TestFeaturizedSamples(unittest.TestCase):
    input_transforms = []
    output_transforms = ["normalize"]
    model_params = {}
    task_types = {"log-solubility": "regression"}
    tasks = ["log-solubility"]
    task_type = "regression"
    task_types = {task: task_type for task in tasks}
    input_file = "../../models/test/example.csv"
    train_samples, valid_samples, test_samples = (
        self._featurize_train_valid_test_split(
            splittype, input_file, task_types.keys(), frac_train=.8,
            splittype, input_file, tasks, frac_train=.8,
            frac_valid=.1, frac_test=.1))
    assert len(train_samples) == 8
    assert len(valid_samples) == 1
@@ -125,11 +131,13 @@ class TestFeaturizedSamples(unittest.TestCase):
    input_transforms = []
    output_transforms = ["normalize"]
    model_params = {}
    task_types = {"log-solubility": "regression"}
    tasks = ["log-solubility"]
    task_type = "regression"
    task_types = {task: task_type for task in tasks}
    input_file = "../../models/test/example.csv"
    train_samples, test_samples = (
        self._featurize_train_valid_test_split(
            splittype, input_file, task_types.keys(), frac_train=.8,
            splittype, input_file, tasks, frac_train=.8,
            frac_valid=0, frac_test=.2))
    assert len(train_samples) == 8
    assert len(test_samples) == 2
+7 −5
Original line number Diff line number Diff line
@@ -15,8 +15,9 @@ class HyperparamOpt(object):
  Provides simple hyperparameter search capabilities.
  """

  def __init__(self, model_class, task_types, fit_transformers=None, verbosity=None):
  def __init__(self, model_class, tasks, task_types, fit_transformers=None, verbosity=None):
    self.model_class = model_class
    self.tasks = tasks
    self.task_types = task_types
    self.fit_transformers = fit_transformers
    assert verbosity in [None, "low", "high"]
@@ -63,11 +64,12 @@ class HyperparamOpt(object):
        model_dir = tempfile.mkdtemp()
      #TODO(JG) Fit transformers for TF models
      if self.fit_transformers:
        model = self.model_class(self.task_types, model_params, model_dir,
                                 fit_transformers=self.fit_transformers,
                                 verbosity=self.verbosity)
        model = self.model_class(
            self.tasks, self.task_types, model_params, model_dir,
            fit_transformers=self.fit_transformers, verbosity=self.verbosity)
      else:
        model = self.model_class(self.task_types, model_params, model_dir,
        model = self.model_class(
            self.tasks, self.task_types, model_params, model_dir,
            verbosity=self.verbosity)
        
      model.fit(train_dataset)
Loading