mirror of
https://github.com/wassname/ray.git
synced 2026-09-09 11:32:43 +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
@@ -202,8 +202,8 @@ class TrialRunnerTest2(unittest.TestCase):
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.RUNNING)
|
||||
self.assertEqual(ray.get(trials[0].runner.set_info.remote(1)), 1)
|
||||
path = runner.trial_executor.save(trials[0])
|
||||
kwargs["restore_path"] = path
|
||||
checkpoint = runner.trial_executor.save(trials[0])
|
||||
kwargs["restore_path"] = checkpoint.value
|
||||
|
||||
runner.add_trial(Trial("__fake", **kwargs))
|
||||
trials = runner.get_trials()
|
||||
@@ -216,7 +216,7 @@ class TrialRunnerTest2(unittest.TestCase):
|
||||
self.assertEqual(trials[0].status, Trial.TERMINATED)
|
||||
self.assertEqual(trials[1].status, Trial.RUNNING)
|
||||
self.assertEqual(ray.get(trials[1].runner.get_info.remote()), 1)
|
||||
self.addCleanup(os.remove, path)
|
||||
self.addCleanup(os.remove, checkpoint.value)
|
||||
|
||||
def testRestoreMetricsAfterCheckpointing(self):
|
||||
ray.init(num_cpus=1, num_gpus=1)
|
||||
@@ -230,9 +230,9 @@ class TrialRunnerTest2(unittest.TestCase):
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.RUNNING)
|
||||
self.assertEqual(ray.get(trials[0].runner.set_info.remote(1)), 1)
|
||||
path = runner.trial_executor.save(trials[0])
|
||||
checkpoint = runner.trial_executor.save(trials[0])
|
||||
runner.trial_executor.stop_trial(trials[0])
|
||||
kwargs["restore_path"] = path
|
||||
kwargs["restore_path"] = checkpoint.value
|
||||
|
||||
runner.add_trial(Trial("__fake", **kwargs))
|
||||
trials = runner.get_trials()
|
||||
@@ -249,7 +249,7 @@ class TrialRunnerTest2(unittest.TestCase):
|
||||
self.assertEqual(trials[1].last_result["timesteps_since_restore"], 20)
|
||||
self.assertEqual(trials[1].last_result["iterations_since_restore"], 2)
|
||||
self.assertGreater(trials[1].last_result["time_since_restore"], 0)
|
||||
self.addCleanup(os.remove, path)
|
||||
self.addCleanup(os.remove, checkpoint.value)
|
||||
|
||||
def testCheckpointingAtEnd(self):
|
||||
ray.init(num_cpus=1, num_gpus=1)
|
||||
|
||||
Reference in New Issue
Block a user