| 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 + "*"))) |