Convert asserts in unittest to pytest (#2529)

This commit is contained in:
Philipp Moritz
2018-08-01 22:32:10 -07:00
committed by Robert Nishihara
parent 9ea57c2a93
commit d8ba667175
10 changed files with 531 additions and 583 deletions
+16 -16
View File
@@ -111,7 +111,7 @@ class TensorFlowTest(unittest.TestCase):
weights[name] += 1.0
variables.set_weights(weights)
self.assertEqual(weights, variables.get_weights())
assert weights == variables.get_weights()
loss2, init2, _, _ = make_linear_network("w", "b")
sess.run(init2)
@@ -123,16 +123,16 @@ class TensorFlowTest(unittest.TestCase):
weights2[name] += 2.0
variables2.set_weights(weights2)
self.assertEqual(weights2, variables2.get_weights())
assert weights2 == variables2.get_weights()
flat_weights = variables2.get_flat() + 2.0
variables2.set_flat(flat_weights)
assert_almost_equal(flat_weights, variables2.get_flat())
variables3 = ray.experimental.TensorFlowVariables([loss2])
self.assertEqual(variables3.sess, None)
assert variables3.sess is None
sess = tf.Session()
variables3.set_session(sess)
self.assertEqual(variables3.sess, sess)
assert variables3.sess == sess
# Test that the variable names for the two different nets are not
# modified by TensorFlow to be unique (i.e., they should already
@@ -153,8 +153,8 @@ class TensorFlowTest(unittest.TestCase):
ray.init(num_workers=1)
net = LossActor(use_loss=False)
self.assertEqual(len(net.values[0].variables.items()), 1)
self.assertEqual(len(net.values[0].placeholders.items()), 1)
assert len(net.values[0].variables.items()) == 1
assert len(net.values[0].placeholders.items()) == 1
net.values[0].set_weights(net.values[0].get_weights())
@@ -164,8 +164,8 @@ class TensorFlowTest(unittest.TestCase):
ray.init(num_workers=1)
net = LossActor()
self.assertEqual(len(net.values[0].variables.items()), 3)
self.assertEqual(len(net.values[0].placeholders.items()), 3)
assert len(net.values[0].variables.items()) == 3
assert len(net.values[0].placeholders.items()) == 3
net.values[0].set_weights(net.values[0].get_weights())
@@ -183,20 +183,20 @@ class TensorFlowTest(unittest.TestCase):
# This only works because at the moment they have size 1.
weights1 = net1.get_weights()
weights2 = net2.get_weights()
self.assertNotEqual(weights1, weights2)
assert weights1 != weights2
# Set the weights and get the weights, and make sure they are
# unchanged.
new_weights1 = net1.set_and_get_weights(weights1)
new_weights2 = net2.set_and_get_weights(weights2)
self.assertEqual(weights1, new_weights1)
self.assertEqual(weights2, new_weights2)
assert weights1 == new_weights1
assert weights2 == new_weights2
# Swap the weights.
new_weights1 = net2.set_and_get_weights(weights1)
new_weights2 = net1.set_and_get_weights(weights2)
self.assertEqual(weights1, new_weights1)
self.assertEqual(weights2, new_weights2)
assert weights1 == new_weights1
assert weights2 == new_weights2
# This test creates an additional network on the driver so that the
# tensorflow variables on the driver and the worker differ.
@@ -214,7 +214,7 @@ class TensorFlowTest(unittest.TestCase):
new_weights2 = ray.get(
net2.set_and_get_weights.remote(net2.get_weights.remote()))
self.assertEqual(weights2, new_weights2)
assert weights2 == new_weights2
def testVariablesControlDependencies(self):
ray.init(num_workers=1)
@@ -228,7 +228,7 @@ class TensorFlowTest(unittest.TestCase):
# Tests if all variables are properly retrieved, 2 variables and 2
# momentum variables.
self.assertEqual(len(net_vars.variables.items()), 4)
assert len(net_vars.variables.items()) == 4
def testRemoteTrainingStep(self):
ray.init(num_workers=1)
@@ -263,7 +263,7 @@ class TensorFlowTest(unittest.TestCase):
sess.run(train, feed_dict=feed_dict)
after_acc = sess.run(
loss, feed_dict=dict(zip(placeholders, [[2] * 100, [4] * 100])))
self.assertTrue(before_acc < after_acc)
assert before_acc < after_acc
if __name__ == "__main__":