mirror of
https://github.com/wassname/ray.git
synced 2026-09-09 11:32:43 +08:00
[tune] option to raise on error (#10030)
This commit is contained in:
@@ -216,6 +216,30 @@ class TrialRunnerTest2(unittest.TestCase):
|
||||
self.assertEqual(trials[0].status, Trial.ERROR)
|
||||
self.assertRaises(TuneError, lambda: runner.step())
|
||||
|
||||
def testFailFastRaise(self):
|
||||
ray.init(num_cpus=1, num_gpus=1)
|
||||
runner = TrialRunner(fail_fast=TrialRunner.RAISE)
|
||||
kwargs = {
|
||||
"resources": Resources(cpu=1, gpu=1),
|
||||
"checkpoint_freq": 1,
|
||||
"max_failures": 0,
|
||||
"config": {
|
||||
"mock_error": True,
|
||||
"persistent_error": True,
|
||||
},
|
||||
}
|
||||
runner.add_trial(Trial("__fake", **kwargs))
|
||||
runner.add_trial(Trial("__fake", **kwargs))
|
||||
trials = runner.get_trials()
|
||||
|
||||
runner.step() # Start trial
|
||||
self.assertEqual(trials[0].status, Trial.RUNNING)
|
||||
runner.step() # Process result, dispatch save
|
||||
self.assertEqual(trials[0].status, Trial.RUNNING)
|
||||
runner.step() # Process save
|
||||
with self.assertRaises(Exception):
|
||||
runner.step() # Error
|
||||
|
||||
def testCheckpointing(self):
|
||||
ray.init(num_cpus=1, num_gpus=1)
|
||||
runner = TrialRunner()
|
||||
|
||||
Reference in New Issue
Block a user