Commit 689a7377 authored by Bharath Ramsundar's avatar Bharath Ramsundar Committed by GitHub
Browse files

Merge pull request #680 from lilleswing/learning-decay

Learning Decay In TensorGraph
parents 899f1ef9 5339fc1f
Loading
Loading
Loading
Loading
+10 −1
Original line number Diff line number Diff line
@@ -437,6 +437,8 @@ class TensorGraph(Model):
      for node in self.topsort():
        node_layer = self.layers[node]
        out_tensors.append(node_layer.none_tensors())
      optimizer = self.optimizer
      self.optimizer = None
      training_placeholder = self._training_placeholder
      self._training_placeholder = None
      self.built = False
@@ -456,6 +458,7 @@ class TensorGraph(Model):
        node_layer = self.layers[node]
        node_layer.set_tensors(out_tensors[index])
      self._training_placeholder = training_placeholder
      self.optimizer = optimizer
      self.built = True
    self.tensor_objects = tensor_objects
    self.rnn_initial_states = rnn_initial_states
@@ -500,6 +503,9 @@ class TensorGraph(Model):
      return tf.get_collection(
          tf.GraphKeys.GLOBAL_VARIABLES, scope=layer.variable_scope)

  def get_global_step(self):
    return self._get_tf("GlobalStep")

  def _get_tf(self, obj):
    """
    TODO(LESWING) REALLY NEED TO DOCUMENT THIS
@@ -523,10 +529,13 @@ class TensorGraph(Model):
      self.tensor_objects['Optimizer'] = self.optimizer()
    elif obj == 'train_op':
      self.tensor_objects['train_op'] = self._get_tf('Optimizer').minimize(
          self.loss.out_tensor)
          self.loss.out_tensor, global_step=self._get_tf('GlobalStep'))
    elif obj == 'summary_op':
      self.tensor_objects['summary_op'] = tf.summary.merge_all(
          key=tf.GraphKeys.SUMMARIES)
    elif obj == 'GlobalStep':
      with self._get_tf("Graph").as_default():
        self.tensor_objects['GlobalStep'] = tf.Variable(0, trainable=False)
    return self._get_tf(obj)

  def _initialize_weights(self, sess, saver):
+35 −1
Original line number Diff line number Diff line
@@ -4,6 +4,7 @@ import numpy as np
import os
from nose.tools import assert_true
from flaky import flaky
import tensorflow as tf

import deepchem as dc
from deepchem.data import NumpyDataset
@@ -11,7 +12,7 @@ from deepchem.data.datasets import Databag
from deepchem.models.tensorgraph.layers import Dense, SoftMaxCrossEntropy, ReduceMean, SoftMax
from deepchem.models.tensorgraph.layers import Feature, Label
from deepchem.models.tensorgraph.layers import ReduceSquareDifference
from deepchem.models.tensorgraph.tensor_graph import TensorGraph
from deepchem.models.tensorgraph.tensor_graph import TensorGraph, TFWrapper


class TestTensorGraph(unittest.TestCase):
@@ -161,6 +162,39 @@ class TestTensorGraph(unittest.TestCase):
    prediction = np.squeeze(tg.predict_proba_on_batch(X))
    assert_true(np.all(np.isclose(prediction, y, atol=0.4)))

  @flaky
  def test_set_optimizer(self):
    n_data_points = 20
    n_features = 2
    X = np.random.rand(n_data_points, n_features)
    y = [[0, 1] for x in range(n_data_points)]
    dataset = NumpyDataset(X, y)
    features = Feature(shape=(None, n_features))
    dense = Dense(out_channels=2, in_layers=[features])
    output = SoftMax(in_layers=[dense])
    label = Label(shape=(None, 2))
    smce = SoftMaxCrossEntropy(in_layers=[label, dense])
    loss = ReduceMean(in_layers=[smce])
    tg = dc.models.TensorGraph(learning_rate=0.01, use_queue=False)
    tg.add_output(output)
    tg.set_loss(loss)
    global_step = tg.get_global_step()

    def optimizer_function():
      starter_learning_rate = 0.1
      learning_rate = tf.train.exponential_decay(
          starter_learning_rate, global_step, 100000, 0.96, staircase=True)
      return tf.train.GradientDescentOptimizer(learning_rate)

    tg.set_optimizer(TFWrapper(optimizer_function))
    tg.fit(dataset, nb_epoch=1000)
    prediction = np.squeeze(tg.predict_proba_on_batch(X))
    tg.save()

    tg1 = TensorGraph.load_from_dir(tg.model_dir)
    prediction2 = np.squeeze(tg1.predict_proba_on_batch(X))
    assert_true(np.all(np.isclose(prediction, prediction2, atol=0.01)))

  def test_tensorboard(self):
    n_data_points = 20
    n_features = 2