Commit 6536f886 authored by leswing's avatar leswing
Browse files

yapf

parent c0687073
Loading
Loading
Loading
Loading
+14 −9
Original line number Diff line number Diff line
@@ -12,6 +12,7 @@ from operator import mul
from deepchem.utils.evaluate import Evaluator
from deepchem.utils.save import log


class HyperparamOpt(object):
  """
  Provides simple hyperparameter search capabilities.
@@ -23,8 +24,13 @@ class HyperparamOpt(object):

  # TODO(rbharath): This function is complicated and monolithic. Is there a nice
  # way to refactor this?
  def hyperparam_search(self, params_dict, train_dataset, valid_dataset,
                        output_transformers, metric, use_max=True,
  def hyperparam_search(self,
                        params_dict,
                        train_dataset,
                        valid_dataset,
                        output_transformers,
                        metric,
                        use_max=True,
                        logdir=None):
    """Perform hyperparams search according to params_dict.
    
@@ -49,10 +55,10 @@ class HyperparamOpt(object):
    best_hyperparams = None
    best_model, best_model_dir = None, None
    all_scores = {}
    for ind, hyperparameter_tuple in enumerate(itertools.product(*hyperparam_vals)):
    for ind, hyperparameter_tuple in enumerate(
        itertools.product(*hyperparam_vals)):
      model_params = {}
      log("Fitting model %d/%d" % (ind+1, number_combinations),
          self.verbose)
      log("Fitting model %d/%d" % (ind + 1, number_combinations), self.verbose)
      for hyperparam, hyperparam_val in zip(hyperparams, hyperparameter_tuple):
        model_params[hyperparam] = hyperparam_val
      log("hyperparameters: %s" % str(model_params), self.verbose)
@@ -92,8 +98,8 @@ class HyperparamOpt(object):
        shutil.rmtree(model_dir)

      log("Model %d/%d, Metric %s, Validation set %s: %f" %
          (ind+1, number_combinations, metric.name, ind, valid_score),
          self.verbose)
          (ind + 1, number_combinations, metric.name, ind,
           valid_score), self.verbose)
      log("\tbest_validation_score so far: %f" % best_validation_score,
          self.verbose)
    if best_model is None:
@@ -107,8 +113,7 @@ class HyperparamOpt(object):
    multitask_scores = train_evaluator.compute_model_performance(
        [metric], train_csv_out.name, train_stats_out.name)
    train_score = multitask_scores[metric.name]
    log("Best hyperparameters: %s" % str(best_hyperparams),
        self.verbose)
    log("Best hyperparameters: %s" % str(best_hyperparams), self.verbose)
    log("train_score: %f" % train_score, self.verbose)
    log("validation_score: %f" % best_validation_score, self.verbose)
    return best_model, best_hyperparams, all_scores