mirror of
https://github.com/wassname/ray.git
synced 2026-09-09 11:32:43 +08:00
[tune] Added timeout parameter to tune.run(), (#10642)
This commit is contained in:
@@ -1108,6 +1108,43 @@ class TrainableFunctionApiTest(unittest.TestCase):
|
||||
self.assertIn("PRINT_STDERR", content)
|
||||
self.assertIn("LOG_STDERR", content)
|
||||
|
||||
def testTimeout(self):
|
||||
from ray.tune.stopper import TimeoutStopper
|
||||
import datetime
|
||||
|
||||
def train(config):
|
||||
for i in range(20):
|
||||
tune.report(metric=i)
|
||||
time.sleep(1)
|
||||
|
||||
register_trainable("f1", train)
|
||||
|
||||
start = time.time()
|
||||
tune.run("f1", time_budget_s=5)
|
||||
diff = time.time() - start
|
||||
self.assertLess(diff, 10)
|
||||
|
||||
# Metric should fire first
|
||||
start = time.time()
|
||||
tune.run("f1", stop={"metric": 3}, time_budget_s=7)
|
||||
diff = time.time() - start
|
||||
self.assertLess(diff, 7)
|
||||
|
||||
# Timeout should fire first
|
||||
start = time.time()
|
||||
tune.run("f1", stop={"metric": 10}, time_budget_s=5)
|
||||
diff = time.time() - start
|
||||
self.assertLess(diff, 10)
|
||||
|
||||
# Combined stopper. Shorter timeout should win.
|
||||
start = time.time()
|
||||
tune.run(
|
||||
"f1",
|
||||
stop=TimeoutStopper(10),
|
||||
time_budget_s=datetime.timedelta(seconds=3))
|
||||
diff = time.time() - start
|
||||
self.assertLess(diff, 9)
|
||||
|
||||
|
||||
class ShimCreationTest(unittest.TestCase):
|
||||
def testCreateScheduler(self):
|
||||
|
||||
Reference in New Issue
Block a user