mirror of
https://github.com/wassname/ray.git
synced 2026-08-20 12:40:44 +08:00
[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:
committed by
Richard Liaw
parent
0834bda8c1
commit
1558307ac4
@@ -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__":
|
||||
|
||||
Reference in New Issue
Block a user