[tune] Prevent MEMORY checkpoints from breaking trial FT (#6691)

* Prevent MEMORY checkpoints from breaking FT

* Add save/pause/resume/restore test

* change checkpoint return value based on status

* Fix test_checkpoint_manager_tests.

* Fix test + checkpoint manager bug

* lint

* Add docstring

* Add docstring to checkpoint_manager constructor

* Change variable name for clarity

* Revert on_checkpoint docstring wording

* Break after success

* nit: more informative warning

* Quarantine test
This commit is contained in:
Ujval Misra
2020-01-22 23:17:09 -08:00
committed by Richard Liaw
parent 0834bda8c1
commit 1558307ac4
10 changed files with 122 additions and 60 deletions
@@ -28,14 +28,14 @@ class CheckpointManagerTest(unittest.TestCase):
for i in range(3)
]
with patch.object(checkpoint_manager, "delete") as \
delete_mock:
with patch.object(checkpoint_manager, "delete") as delete_mock:
for j in range(3):
checkpoint_manager.on_checkpoint(checkpoints[j])
expected_deletes = 0 if j != 2 else 1
self.assertEqual(delete_mock.call_count, expected_deletes, j)
self.assertEqual(checkpoint_manager.newest_checkpoint,
checkpoints[j])
self.assertEqual(
checkpoint_manager.newest_persistent_checkpoint,
checkpoints[j])
best_checkpoints = checkpoint_manager.best_checkpoints()
self.assertEqual(len(best_checkpoints), keep_checkpoints_num)
@@ -59,8 +59,9 @@ class CheckpointManagerTest(unittest.TestCase):
checkpoint_manager.on_checkpoint(checkpoints[j])
expected_deletes = 0 if j != 3 else 1
self.assertEqual(delete_mock.call_count, expected_deletes)
self.assertEqual(checkpoint_manager.newest_checkpoint,
checkpoints[j])
self.assertEqual(
checkpoint_manager.newest_persistent_checkpoint,
checkpoints[j])
best_checkpoints = checkpoint_manager.best_checkpoints()
self.assertEqual(len(best_checkpoints), keep_checkpoints_num)
@@ -74,7 +75,7 @@ class CheckpointManagerTest(unittest.TestCase):
keep_checkpoints_num = 4
checkpoint_manager = self.checkpoint_manager(keep_checkpoints_num)
checkpoints = [
Checkpoint(Checkpoint.MEMORY, i, self.mock_result(i))
Checkpoint(Checkpoint.PERSISTENT, i, self.mock_result(i))
for i in range(16)
]
random.shuffle(checkpoints)
@@ -92,15 +93,28 @@ class CheckpointManagerTest(unittest.TestCase):
Tests that an error is logged when the associated result of the
checkpoint has no checkpoint score attribute.
"""
keep_checkpoints_num = 1
checkpoint_manager = self.checkpoint_manager(keep_checkpoints_num)
checkpoint_manager = self.checkpoint_manager(keep_checkpoints_num=1)
no_attr_checkpoint = Checkpoint(Checkpoint.MEMORY, 0, {})
no_attr_checkpoint = Checkpoint(Checkpoint.PERSISTENT, 0, {})
with patch.object(logger, "error") as log_error_mock:
checkpoint_manager.on_checkpoint(no_attr_checkpoint)
log_error_mock.assert_called_once()
# The newest checkpoint should still be set despite this error.
assert checkpoint_manager.newest_checkpoint == no_attr_checkpoint
self.assertEqual(checkpoint_manager.newest_persistent_checkpoint,
no_attr_checkpoint)
def testOnMemoryCheckpoint(self):
checkpoints = [
Checkpoint(Checkpoint.MEMORY, 0, self.mock_result(0)),
Checkpoint(Checkpoint.MEMORY, 0, self.mock_result(0))
]
checkpoint_manager = self.checkpoint_manager(keep_checkpoints_num=1)
checkpoint_manager.on_checkpoint(checkpoints[0])
checkpoint_manager.on_checkpoint(checkpoints[1])
newest = checkpoint_manager.newest_memory_checkpoint
self.assertEqual(newest, checkpoints[1])
self.assertEqual(checkpoint_manager.best_checkpoints(), [])
if __name__ == "__main__":