MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / testSharded

Method testSharded

tensorflow/python/training/saver_test.py:1384–1443  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

1382 save_dir=save_dir)
1383
1384 def testSharded(self):
1385 save_dir = self._get_test_dir("max_to_keep_sharded")
1386
1387 with session.Session(
1388 target="",
1389 config=config_pb2.ConfigProto(device_count={"CPU": 2})) as sess:
1390 with sess.graph.device("/cpu:0"):
1391 v0 = variables.VariableV1(111, name="v0")
1392 with sess.graph.device("/cpu:1"):
1393 v1 = variables.VariableV1(222, name="v1")
1394 save = saver_module.Saver(
1395 {
1396 "v0": v0,
1397 "v1": v1
1398 }, sharded=True, max_to_keep=2)
1399 self.evaluate(variables.global_variables_initializer())
1400 self.assertEqual([], save.last_checkpoints)
1401
1402 s1 = save.save(sess, os.path.join(save_dir, "s1"))
1403 self.assertEqual([s1], save.last_checkpoints)
1404 if save._write_version is saver_pb2.SaverDef.V1:
1405 self.assertEqual(2, len(gfile.Glob(s1)))
1406 else:
1407 self.assertEqual(4, len(gfile.Glob(s1 + "*")))
1408
1409 self.assertTrue(
1410 gfile.Exists(checkpoint_management.meta_graph_filename(s1)))
1411
1412 s2 = save.save(sess, os.path.join(save_dir, "s2"))
1413 self.assertEqual([s1, s2], save.last_checkpoints)
1414 if save._write_version is saver_pb2.SaverDef.V1:
1415 self.assertEqual(2, len(gfile.Glob(s1)))
1416 else:
1417 self.assertEqual(4, len(gfile.Glob(s1 + "*")))
1418 self.assertTrue(
1419 gfile.Exists(checkpoint_management.meta_graph_filename(s1)))
1420 if save._write_version is saver_pb2.SaverDef.V1:
1421 self.assertEqual(2, len(gfile.Glob(s2)))
1422 else:
1423 self.assertEqual(4, len(gfile.Glob(s2 + "*")))
1424 self.assertTrue(
1425 gfile.Exists(checkpoint_management.meta_graph_filename(s2)))
1426
1427 s3 = save.save(sess, os.path.join(save_dir, "s3"))
1428 self.assertEqual([s2, s3], save.last_checkpoints)
1429 self.assertEqual(0, len(gfile.Glob(s1 + "*")))
1430 self.assertFalse(
1431 gfile.Exists(checkpoint_management.meta_graph_filename(s1)))
1432 if save._write_version is saver_pb2.SaverDef.V1:
1433 self.assertEqual(2, len(gfile.Glob(s2)))
1434 else:
1435 self.assertEqual(4, len(gfile.Glob(s2 + "*")))
1436 self.assertTrue(
1437 gfile.Exists(checkpoint_management.meta_graph_filename(s2)))
1438 if save._write_version is saver_pb2.SaverDef.V1:
1439 self.assertEqual(2, len(gfile.Glob(s3)))
1440 else:
1441 self.assertEqual(4, len(gfile.Glob(s3 + "*")))

Callers

nothing calls this directly

Calls 7

_get_test_dirMethod · 0.95
saveMethod · 0.95
SessionMethod · 0.45
deviceMethod · 0.45
evaluateMethod · 0.45
joinMethod · 0.45
ExistsMethod · 0.45

Tested by

no test coverage detected