Commit 57db5d78 authored by Bharath Ramsundar's avatar Bharath Ramsundar
Browse files

Some bugfixes

parent 9bb5df8f
Loading
Loading
Loading
Loading
+3 −3
Original line number Diff line number Diff line
@@ -19,13 +19,13 @@ class KerasModel(Model):
  Abstract base class shared across all Keras models.
  """

  def save(self, out_dir):
  def save(self):
    """
    Saves underlying keras model to disk.
    """
    super(KerasModel, self).save(out_dir)
    super(KerasModel, self).save()
    model = self.get_raw_model()
    filename, _ = os.path.splitext(Model.get_model_filename(out_dir))
    filename, _ = os.path.splitext(Model.get_model_filename(self.model_dir))

    # Note that keras requires the model architecture and weights to be stored
    # separately. A json file is generated that specifies the model architecture.
+3 −2
Original line number Diff line number Diff line
@@ -140,8 +140,9 @@ class SingleTaskDNN(MultiTaskDNN):
  """
  Abstract base class for different ML models.
  """
  def __init__(self, task_types, model_params, fit_transformers=None, initialize_raw_model=True, verbosity="low"):
    super(SingleTaskDNN, self).__init__(task_types, model_params,
  def __init__(self, task_types, model_params, model_dir, fit_transformers=None,
               initialize_raw_model=True, verbosity="low"):
    super(SingleTaskDNN, self).__init__(task_types, model_params, model_dir,
                                        fit_transformers=fit_transformers,
                                        initialize_raw_model=initialize_raw_model,
                                        verbosity=verbosity)
+2 −4
Original line number Diff line number Diff line
@@ -729,13 +729,11 @@ class TensorflowModel(Model):
    """
    return self.eval_model.predict_on_batch(X)

  def save(self, logdir):
  def save(self):
    """
    No-op since tf models save themselves during fit()
    """
    if logdir != self.train_model.logdir:
      raise ValueError("Cannot save to directory "
                       "that was not specified during initialization")
    pass

  def load(self, model_dir):
    """
+1 −1
Original line number Diff line number Diff line
@@ -73,7 +73,7 @@ class TestAPI(unittest.TestCase):

    # Fit trained model
    model.fit(train_dataset)
    model.save(self.model_dir)
    model.save()

    # Eval model on train
    evaluator = Evaluator(model, train_dataset, transformers, verbose=True)
+6 −6
Original line number Diff line number Diff line
@@ -54,7 +54,7 @@ class TestKerasSklearnAPI(TestAPI):
                          Metric(metrics.mean_squared_error),
                          Metric(metrics.mean_absolute_error)]

    model = SklearnModel(task_types, model_params,
    model = SklearnModel(task_types, model_params, self.model_dir,
                         model_instance=RandomForestRegressor())
    self._create_model(train_dataset, test_dataset, model, transformers,
                       regression_metrics)
@@ -82,7 +82,7 @@ class TestKerasSklearnAPI(TestAPI):
                          Metric(metrics.mean_squared_error),
                          Metric(metrics.mean_absolute_error)]

    model = SklearnModel(task_types, model_params,
    model = SklearnModel(task_types, model_params, self.model_dir,
                         model_instance=RandomForestRegressor())
    self._create_model(train_dataset, test_dataset, model, transformers,
                       regression_metrics)
@@ -109,7 +109,7 @@ class TestKerasSklearnAPI(TestAPI):
                          Metric(metrics.mean_squared_error),
                          Metric(metrics.mean_absolute_error)]

    model = SklearnModel(task_types, model_params,
    model = SklearnModel(task_types, model_params, self.model_dir,
                         model_instance=RandomForestRegressor())
    self._create_model(train_dataset, test_dataset, model, transformers,
                       regression_metrics)
@@ -133,7 +133,7 @@ class TestKerasSklearnAPI(TestAPI):
                          Metric(metrics.mean_squared_error),
                          Metric(metrics.mean_absolute_error)]

    model = SklearnModel(task_types, model_params,
    model = SklearnModel(task_types, model_params, self.model_dir,
                         model_instance=RandomForestRegressor())
    self._create_model(train_dataset, test_dataset, model, transformers,
                       regression_metrics)
@@ -206,7 +206,7 @@ class TestKerasSklearnAPI(TestAPI):
                          Metric(metrics.mean_squared_error),
                          Metric(metrics.mean_absolute_error)]

    model = SingleTaskDNN(task_types, model_params)
    model = SingleTaskDNN(task_types, model_params, self.model_dir)
    self._create_model(train_dataset, test_dataset, model, transformers,
                       regression_metrics)

@@ -281,6 +281,6 @@ class TestKerasSklearnAPI(TestAPI):
                              Metric(metrics.recall_score),
                              Metric(metrics.accuracy_score)]
    
    model = MultiTaskDNN(task_types, model_params)
    model = MultiTaskDNN(task_types, model_params, self.model_dir)
    self._create_model(train_dataset, test_dataset, model, transformers,
                       classification_metrics)