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

Removing some unneeded model_config files

parent 36d8d142
Loading
Loading
Loading
Loading
+0 −211
Original line number Diff line number Diff line
@@ -184,214 +184,3 @@ class TensorflowMultiTaskClassifier(TensorflowClassifier):
    for name, value in named_values.iteritems():
      feed_dict['{}/{}:0'.format(self.placeholder_root, name)] = value
    return feed_dict

  def ReadInput(self, input_pattern, input_data_types=None):
    """Read input data and return a generator for minibatches.

    Args:
      input_pattern: Input file pattern.
      input_data_types: List of legacy_types_pb2 constants matching the
          number of and data types present in the sstables. If not specified,
          defaults to full ICML 259-task types, but can be specified
          for unittests or other datasets with consistent types.

    Returns:
      A generator that yields a dict for feeding a single batch to Placeholders
      in the graph.

    Raises:
      AssertionError: If no default session is available.
    """
    if model_ops.IsTraining():
      randomize = True
      num_iterations = None
    else:
      randomize = False
      num_iterations = 1

    num_tasks = self.model_params["num_classification_tasks"]
    tasks_in_input = self.model_params["tasks_in_input"]
    if input_data_types is None:
      input_data_types = ([legacy_types_pb2.DF_FLOAT] +
                          [legacy_types_pb2.DF_LABEL_PROTO] * tasks_in_input)
    features, labels = input_ops.InputExampleInputReader(
        input_pattern=input_pattern,
        batch_size=self.model_params["batch_size"],
        num_tasks=num_tasks,
        input_data_types=input_data_types,
        num_features=self.model_params["num_features"],
        randomize=randomize,
        shuffling=randomize,
        num_iterations=num_iterations)

    return self._ReadInputGenerator(features, labels[:, :num_tasks])

  def _GetFeedDict(self, named_values):
    feed_dict = {}
    for name, value in named_values.iteritems():
      feed_dict['{}/{}:0'.format(self.placeholder_root, name)] = value

    return feed_dict

  def EvalBatch(self, input_batch):
    """Runs inference on the provided batch of input.

    Args:
      input_batch: iterator of input with len self.model_params["batch_size"].

    Returns:
      Tuple of three numpy arrays with shape num_examples x num_tasks (x ...):
        output: Model predictions.
        labels: True labels. numpy array values are scalars,
            not 1-hot classes vector.
        weights: Example weights.
    """
    output, labels, weights = super(TensorflowMultiTaskDNN, self).EvalBatch(
        input_batch)

    # Converts labels from 1-hot to float.
    labels = labels[:, :, 1]  # Whole batch, all tasks, 1-hot positive index.
    return output, labels, weights

  def BatchInputGenerator(self, serialized_batch):
    """Returns a generator that iterates over the provided batch of input.

    TODO(user): This is similar to input_ops.InputExampleInputReader(),
        but doesn't need to be executed as part of the TensorFlow graph.
        Consider refactoring so these can share code somehow.

    Args:
      serialized_batch: List of tuples: (_, value) where value is
          a serialized InputExample proto. Must have self.model_params["batch_size"]
          length or smaller. If smaller, we'll pad up to batch_size
          and mark the padding as invalid so it's ignored in eval metrics.
    Yields:
      Dict of model inputs for use as a feed_dict.

    Raises:
      ValueError: If the batch is larger than the batch_size.
    """
    if len(serialized_batch) > self.model_params["batch_size"]:
      raise ValueError(
          'serialized_batch length {} must be <= batch_size {}'.format(
              len(serialized_batch), self.model_params["batch_size"]))
    for _ in xrange(self.model_params["batch_size"] - len(serialized_batch)):
      serialized_batch.append((None, ''))

    features = []
    labels = []
    for _, serialized_proto in serialized_batch:
      if serialized_proto:
        input_example = input_example_pb2.InputExample()
        input_example.ParseFromString(serialized_proto)
        features.append([f for f in input_example.endpoint[0].float_value])
        label_protos = [endpoint.label
                        for endpoint in input_example.endpoint[1:]]
        assert len(label_protos) == self.model_params["num_classification_tasks"]
        labels.append([l.SerializeToString() for l in label_protos])
      else:
        # This was a padded value to reach the batch size.
        features.append([0.0 for _ in xrange(self.model_params["num_features"])])
        labels.append(
            ['' for _ in xrange(self.model_params["num_classification_tasks"])])

    valid = np.asarray([(np.sum(f) > 0) for f in features])

    assert len(features) == self.model_params["batch_size"]
    assert len(labels) == self.model_params["batch_size"]
    assert len(valid) == self.model_params["batch_size"]
    yield self._GetFeedDict({
        'mol_features': features,
        'labels': labels,
        'valid': valid
    })

  def _ReadInputGenerator(self, features_tensor, labels_tensor):
    """Generator that constructs feed_dict for minibatches.

    Args:
      features_tensor: Tensor of batch_size x molecule features.
      labels_tensor: Tensor of batch_size x label protos.

    Yields:
      A dict for feeding a single batch to Placeholders in the graph.

    Raises:
      AssertionError: If no default session is available.
    """
    sess = tf.get_default_session()
    if sess is None:
      raise AssertionError('No default session')
    while True:
      try:
        logging.vlog(1, 'Starting session execution to get input data')
        features, labels = sess.run([features_tensor, labels_tensor])
        logging.vlog(1, 'Done with session execution to get input data')
        # TODO(user): check if the below axis=1 needs to change to axis=0,
        # because cl/105081140.
        valid = np.sum(features, axis=1) > 0
        yield self._GetFeedDict({
            'mol_features': features,
            'labels': labels,
            'valid': valid
        })

      except tf.OpError as e:
        # InputExampleInput op raises OpError when it has hit num_iterations
        # or its input file is exhausted. However it may also be raised
        # if the input sstable isn't what we expect.
        if 'Invalid InputExample' in e.message:
          raise e
        else:
          break

  def Run(self, input_data_types=None):
    """Trains the model with specified parameters.

    Args:
      input_data_types: List of legacy_types_pb2 constants or None.
    """
    model_params = model_config.ModelConfig({
        'input_pattern': '',  # Should have %d for fold index substitution.
        'num_classification_tasks': 259,
        'tasks_in_input': 259,  # Dimensionality of sstables
        'max_steps': 50000000,
        'summaries': False,
        'batch_size': 128,
        'learning_rate': 0.0003,
        'num_classes': 2,
        'optimizer': 'sgd',
        'penalty': 0.0,
        'num_features': 1024,
        'layer_sizes': [1200],
        'weight_init_stddevs': [0.01],
        'bias_init_consts': [0.5],
        'dropouts': [0.0],
    })
    #model_params.ReadFromFile(FLAGS.config,
    #                          overwrite='required')

    if FLAGS.replica_id == 0:
      gfile.MakeDirs(FLAGS.logdir)
      #model_params.WriteToFile(os.path.join(FLAGS.logdir, 'config.pbtxt'))

#    model = icml_models.IcmlModel(config,
#                                  train=True,
#                                  logdir=FLAGS.logdir,
#                                  master=FLAGS.master)

    if FLAGS.num_folds is not None and FLAGS.fold is not None:
      folds = kfold_pattern(config.input_pattern, FLAGS.num_folds,
                            FLAGS.fold)
      train_pattern, _ = folds.next()
      train_pattern = ','.join(train_pattern)
    else:
      train_pattern = config.input_pattern

    with model.graph.as_default():
      model.fit(model.read_input(train_pattern,
                                 input_data_types=input_data_types),
                max_steps=config.max_steps,
                summaries=config.summaries,
                replica_id=FLAGS.replica_id,
                ps_tasks=FLAGS.ps_tasks)
+0 −78
Original line number Diff line number Diff line
#!/usr/bin/python
#
# Copyright 2015 Google Inc.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#      http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Evaluate a model from the ICML-2015 paper.

This script requires a trained model with its associated config and checkpoint.
If you don't have a trained model, run icml_train.py first.

"""
# pylint: disable=line-too-long
# pylint: enable=line-too-long


from nowhere.research.biology.collaborations.pande.py import utils

from tensorflow.python.platform import app
from tensorflow.python.platform import flags
from tensorflow.python.platform import gfile

from biology import model_config
from biology.icml import icml_models

flags.DEFINE_string('config', None, 'Serialized ModelConfig proto.')
flags.DEFINE_string('checkpoint', None,
                    'Model checkpoint file. File can contain either an '
                    'absolute checkpoint (e.g. model.ckpt-{step}) or a '
                    'serialized CheckpointState proto.')
flags.DEFINE_string('input_pattern', None, 'Input file pattern; '
                    'It should include %d for fold index substitution.')
flags.DEFINE_string('master', 'local', 'BNS name of the TensorFlow master.')
flags.DEFINE_string('logdir', None, 'Directory for output files.')
flags.DEFINE_integer('num_folds', 5, 'Number of cross-validation folds.')
flags.DEFINE_integer('fold', None, 'Fold index for this model.')
flags.DEFINE_enum('model_type', 'single', ['single', 'deep', 'deepaux', 'py',
                                           'pydrop1', 'pydrop2'],
                  'Which model from the ICML paper should be trained/evaluated')
FLAGS = flags.FLAGS


def main(unused_argv=None):
  config = model_config.ModelConfig()
  config.ReadFromFile(FLAGS.config, overwrite='allowed')
  gfile.MakeDirs(FLAGS.logdir)
  model = icml_models.CONSTRUCTORS[FLAGS.model_type](config,
                                                     train=False,
                                                     logdir=FLAGS.logdir,
                                                     master=FLAGS.master)

  if FLAGS.num_folds is not None and FLAGS.fold is not None:
    folds = utils.kfold_pattern(FLAGS.input_pattern, FLAGS.num_folds,
                                FLAGS.fold)
    _, test_pattern = folds.next()
    test_pattern = ','.join(test_pattern)
  else:
    test_pattern = FLAGS.input_pattern

  with model.graph.as_default():
    model.Eval(model.ReadInput(test_pattern), FLAGS.checkpoint)


if __name__ == '__main__':
  flags.MarkFlagAsRequired('config')
  flags.MarkFlagAsRequired('checkpoint')
  flags.MarkFlagAsRequired('input_pattern')
  flags.MarkFlagAsRequired('logdir')
  app.run()
+0 −0

Empty file deleted.

+0 −81
Original line number Diff line number Diff line
#!/usr/bin/python
#
# Copyright 2015 Google Inc.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#      http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Train a model from the ICML-2015 paper.
"""
# pylint: disable=line-too-long
# pylint: enable=line-too-long

import os


from tensorflow.python.platform import app
from tensorflow.python.platform import flags
from tensorflow.python.platform import gfile

from biology import model_config
from biology.icml import icml_models

flags.DEFINE_string('config', None, 'Serialized ModelConfig proto.')
flags.DEFINE_string('master', '', 'BNS name of the TensorFlow master.')
flags.DEFINE_string('logdir', None, 'Directory for output files.')
flags.DEFINE_integer('replica_id', 0, 'Task ID of this replica.')
flags.DEFINE_integer('ps_tasks', 0, 'Number of parameter server tasks.')
flags.DEFINE_integer('num_folds', 5, 'Number of cross-validation folds.')
flags.DEFINE_integer('fold', None, 'Fold index for this model.')

FLAGS = flags.FLAGS

def kfold_pattern(input_pattern, num_folds, fold=None):
  """Generator for train/test filename splits.

  The pattern is not expanded except for the %d being replaced by the fold
  index.

  Args:
    input_pattern: Input filename pattern. Should contain %d for fold index.
    num_folds: Number of folds.
    fold: If not None, the generator only yields the train/test split for the
      given fold.

  Yields:
    train_filenames: A list of file patterns in training set.
    test_filenames: A list of file patterns in test set.
  """
  # get filenames associated with each fold
  fold_filepatterns = [input_pattern % i for i in range(num_folds)]

  # create train/test splits
  for i in range(num_folds):
    if fold is not None and i != fold:
      continue
    train = fold_filepatterns[:i] + fold_filepatterns[i+1:]
    test = [fold_filepatterns[i]]
    if any([f in test for f in train]):
      logging.fatal('Train/test split is not complete.')
    if set(train + test) != set(fold_filepatterns):
      logging.fatal('Not all input files are accounted for.')
    yield train, test


def main(unused_argv=None):
  Run()


if __name__ == '__main__':
  flags.MarkFlagAsRequired('config')
  flags.MarkFlagAsRequired('logdir')
  flags.MarkFlagAsRequired('fold')
  app.run()
+0 −44
Original line number Diff line number Diff line
// Copyright 2015 Google Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
//      http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//
////////////////////////////////////////////////////////////////////////////////
syntax = "proto2";

package deepchem.models.tensorflow_models;

// Neural network model configuration, used mostly for
// de/serializing model parameters for a given execution from/to disk.
message ModelConfig {
  message Parameter {
    optional string name = 1;
    optional string description = 10;

    // See oneof user guide:
    // http://sites/protocol-buffers/user-docs/miscellaneous-howtos/oneof
    oneof value {
      float float_value = 2;
      int32 int_value = 3;
      string string_value = 4;
      bool bool_value = 5;
    }

    repeated float float_list = 6;
    repeated int32 int_list = 7;
    repeated string string_list = 8;
    repeated bool bool_list = 9;
  };

  repeated Parameter parameter = 1;
  optional string description = 2;
}
Loading