mirror of
https://github.com/wassname/ray.git
synced 2026-08-13 12:30:18 +08:00
[tune] Experiment Management API (#1328)
* init for exposing external interface * revisions * http server * small * simplify * ui * fixes * test * nit * nit * merge * untested * nits * nit * init tests * tests * more tests * nit * fix hyperband * cleanup * nits * good stuff * cleanup * comments and need to test * nit * notebook * testing * test and expose server * server_tests * docs * periods * fix tests * committing test * fi
This commit is contained in:
@@ -497,6 +497,46 @@ class TrialRunnerTest(unittest.TestCase):
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.TERMINATED)
|
||||
|
||||
def testStopTrial(self):
|
||||
ray.init(num_cpus=4, num_gpus=2)
|
||||
runner = TrialRunner()
|
||||
kwargs = {
|
||||
"stopping_criterion": {"training_iteration": 5},
|
||||
"resources": Resources(cpu=1, gpu=1),
|
||||
}
|
||||
trials = [
|
||||
Trial("__fake", **kwargs),
|
||||
Trial("__fake", **kwargs),
|
||||
Trial("__fake", **kwargs),
|
||||
Trial("__fake", **kwargs)]
|
||||
for t in trials:
|
||||
runner.add_trial(t)
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.RUNNING)
|
||||
self.assertEqual(trials[1].status, Trial.PENDING)
|
||||
|
||||
# Stop trial while running
|
||||
runner.stop_trial(trials[0])
|
||||
self.assertEqual(trials[0].status, Trial.TERMINATED)
|
||||
self.assertEqual(trials[1].status, Trial.PENDING)
|
||||
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.TERMINATED)
|
||||
self.assertEqual(trials[1].status, Trial.RUNNING)
|
||||
self.assertEqual(trials[-1].status, Trial.PENDING)
|
||||
|
||||
# Stop trial while pending
|
||||
runner.stop_trial(trials[-1])
|
||||
self.assertEqual(trials[0].status, Trial.TERMINATED)
|
||||
self.assertEqual(trials[1].status, Trial.RUNNING)
|
||||
self.assertEqual(trials[-1].status, Trial.TERMINATED)
|
||||
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.TERMINATED)
|
||||
self.assertEqual(trials[1].status, Trial.RUNNING)
|
||||
self.assertEqual(trials[2].status, Trial.RUNNING)
|
||||
self.assertEqual(trials[-1].status, Trial.TERMINATED)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main(verbosity=2)
|
||||
|
||||
@@ -140,8 +140,25 @@ class EarlyStoppingSuite(unittest.TestCase):
|
||||
|
||||
|
||||
class _MockTrialRunner():
|
||||
def _stop_trial(self, trial):
|
||||
trial.stop()
|
||||
def __init__(self, scheduler):
|
||||
self._scheduler_alg = scheduler
|
||||
|
||||
def process_action(self, trial, action):
|
||||
if action == TrialScheduler.CONTINUE:
|
||||
pass
|
||||
elif action == TrialScheduler.PAUSE:
|
||||
self._pause_trial(trial)
|
||||
elif action == TrialScheduler.STOP:
|
||||
trial.stop()
|
||||
|
||||
def stop_trial(self, trial):
|
||||
if trial.status in [Trial.ERROR, Trial.TERMINATED]:
|
||||
return
|
||||
elif trial.status in [Trial.PENDING, Trial.PAUSED]:
|
||||
self._scheduler_alg.on_trial_remove(self, trial)
|
||||
else:
|
||||
|
||||
self._scheduler_alg.on_trial_complete(self, trial, result(100, 10))
|
||||
|
||||
def has_resources(self, resources):
|
||||
return True
|
||||
@@ -168,7 +185,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
for i in range(num_trials):
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
runner = _MockTrialRunner()
|
||||
runner = _MockTrialRunner(sched)
|
||||
return sched, runner
|
||||
|
||||
def default_statistics(self):
|
||||
@@ -186,14 +203,6 @@ class HyperbandSuite(unittest.TestCase):
|
||||
def downscale(self, n, sched):
|
||||
return int(np.ceil(n / sched._eta))
|
||||
|
||||
def process(self, trl, mock_runner, action):
|
||||
if action == TrialScheduler.CONTINUE:
|
||||
pass
|
||||
elif action == TrialScheduler.PAUSE:
|
||||
mock_runner._pause_trial(trl)
|
||||
elif action == TrialScheduler.STOP:
|
||||
self.stopTrial(trl, mock_runner)
|
||||
|
||||
def basicSetup(self):
|
||||
"""Setup and verify full band.
|
||||
"""
|
||||
@@ -224,10 +233,6 @@ class HyperbandSuite(unittest.TestCase):
|
||||
|
||||
return sched
|
||||
|
||||
def stopTrial(self, trial, mock_runner):
|
||||
self.assertNotEqual(trial.status, Trial.TERMINATED)
|
||||
mock_runner._stop_trial(trial)
|
||||
|
||||
def testConfigSameEta(self):
|
||||
sched = HyperBandScheduler()
|
||||
i = 0
|
||||
@@ -283,7 +288,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
mock_runner, trl, result(cur_units, i))
|
||||
if i < current_length - 1:
|
||||
self.assertEqual(action, TrialScheduler.PAUSE)
|
||||
self.process(trl, mock_runner, action)
|
||||
mock_runner.process_action(trl, action)
|
||||
|
||||
self.assertEqual(action, TrialScheduler.CONTINUE)
|
||||
new_length = len(big_bracket.current_trials())
|
||||
@@ -304,7 +309,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
for i, trl in reversed(list(enumerate(big_bracket.current_trials()))):
|
||||
action = sched.on_trial_result(
|
||||
mock_runner, trl, result(cur_units, i))
|
||||
self.process(trl, mock_runner, action)
|
||||
mock_runner.process_action(trl, action)
|
||||
|
||||
self.assertEqual(action, TrialScheduler.STOP)
|
||||
|
||||
@@ -321,7 +326,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
for i, trl in enumerate(big_bracket.current_trials()):
|
||||
action = sched.on_trial_result(
|
||||
mock_runner, trl, result(cur_units, i))
|
||||
self.process(trl, mock_runner, action)
|
||||
mock_runner.process_action(trl, action)
|
||||
|
||||
self.assertEqual(action, TrialScheduler.CONTINUE)
|
||||
|
||||
@@ -412,9 +417,9 @@ class HyperbandSuite(unittest.TestCase):
|
||||
mock_runner._launch_trial(t)
|
||||
|
||||
for i, t in enumerate(bracket_trials):
|
||||
status = sched.on_trial_result(
|
||||
action = sched.on_trial_result(
|
||||
mock_runner, t, result(init_units, i))
|
||||
self.assertEqual(status, TrialScheduler.CONTINUE)
|
||||
self.assertEqual(action, TrialScheduler.CONTINUE)
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
mock_runner._launch_trial(t)
|
||||
@@ -442,7 +447,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
for i in range(stats["max_trials"]):
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
runner = _MockTrialRunner()
|
||||
runner = _MockTrialRunner(sched)
|
||||
|
||||
big_bracket = sched._hyperbands[0][-1]
|
||||
|
||||
@@ -452,17 +457,11 @@ class HyperbandSuite(unittest.TestCase):
|
||||
|
||||
# Provides results from 0 to 8 in order, keeping the last one running
|
||||
for i, trl in enumerate(big_bracket.current_trials()):
|
||||
status = sched.on_trial_result(runner, trl, result2(1, i))
|
||||
if status == TrialScheduler.CONTINUE:
|
||||
continue
|
||||
elif status == TrialScheduler.PAUSE:
|
||||
runner._pause_trial(trl)
|
||||
elif status == TrialScheduler.STOP:
|
||||
self.assertNotEqual(trl.status, Trial.TERMINATED)
|
||||
self.stopTrial(trl, runner)
|
||||
action = sched.on_trial_result(runner, trl, result2(1, i))
|
||||
runner.process_action(trl, action)
|
||||
|
||||
new_length = len(big_bracket.current_trials())
|
||||
self.assertEqual(status, TrialScheduler.CONTINUE)
|
||||
self.assertEqual(action, TrialScheduler.CONTINUE)
|
||||
self.assertEqual(new_length, self.downscale(current_length, sched))
|
||||
|
||||
def testJumpingTime(self):
|
||||
@@ -476,21 +475,38 @@ class HyperbandSuite(unittest.TestCase):
|
||||
main_trials = big_bracket.current_trials()[:-1]
|
||||
jump = big_bracket.current_trials()[-1]
|
||||
for i, trl in enumerate(main_trials):
|
||||
status = sched.on_trial_result(mock_runner, trl, result(1, i))
|
||||
if status == TrialScheduler.CONTINUE:
|
||||
continue
|
||||
elif status == TrialScheduler.PAUSE:
|
||||
mock_runner._pause_trial(trl)
|
||||
elif status == TrialScheduler.STOP:
|
||||
self.assertNotEqual(trl.status, Trial.TERMINATED)
|
||||
self.stopTrial(trl, mock_runner)
|
||||
action = sched.on_trial_result(mock_runner, trl, result(1, i))
|
||||
mock_runner.process_action(trl, action)
|
||||
|
||||
status = sched.on_trial_result(mock_runner, jump, result(4, i))
|
||||
self.assertEqual(status, TrialScheduler.PAUSE)
|
||||
action = sched.on_trial_result(mock_runner, jump, result(4, i))
|
||||
self.assertEqual(action, TrialScheduler.PAUSE)
|
||||
|
||||
current_length = len(big_bracket.current_trials())
|
||||
self.assertLess(current_length, 27)
|
||||
|
||||
def testRemove(self):
|
||||
"""Test with 4: start 1, remove 1 pending, add 2, remove 1 pending"""
|
||||
sched, runner = self.schedulerSetup(4)
|
||||
trials = sorted(list(sched._trial_info), key=lambda t: t.trial_id)
|
||||
runner._launch_trial(trials[0])
|
||||
sched.on_trial_result(runner, trials[0], result(1, 5))
|
||||
self.assertEqual(trials[0].status, Trial.RUNNING)
|
||||
self.assertEqual(trials[1].status, Trial.PENDING)
|
||||
|
||||
bracket, _ = sched._trial_info[trials[1]]
|
||||
self.assertTrue(trials[1] in bracket._live_trials)
|
||||
sched.on_trial_remove(runner, trials[1])
|
||||
self.assertFalse(trials[1] in bracket._live_trials)
|
||||
|
||||
for i in range(2):
|
||||
trial = Trial("__fake")
|
||||
sched.on_trial_add(None, trial)
|
||||
|
||||
bracket, _ = sched._trial_info[trial]
|
||||
self.assertTrue(trial in bracket._live_trials)
|
||||
sched.on_trial_remove(runner, trial) # where trial is not running
|
||||
self.assertFalse(trial in bracket._live_trials)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main(verbosity=2)
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import unittest
|
||||
import socket
|
||||
|
||||
import ray
|
||||
from ray.rllib import _register_all
|
||||
from ray.tune.trial import Trial, Resources
|
||||
from ray.tune.web_server import TuneClient
|
||||
from ray.tune.trial_runner import TrialRunner
|
||||
|
||||
|
||||
def get_valid_port():
|
||||
port = 4321
|
||||
while True:
|
||||
try:
|
||||
print("Trying port", port)
|
||||
port_test_socket = socket.socket()
|
||||
port_test_socket.bind(("127.0.0.1", port))
|
||||
port_test_socket.close()
|
||||
break
|
||||
except socket.error:
|
||||
port += 1
|
||||
return port
|
||||
|
||||
|
||||
class TuneServerSuite(unittest.TestCase):
|
||||
def basicSetup(self):
|
||||
ray.init(num_cpus=4, num_gpus=1)
|
||||
port = get_valid_port()
|
||||
self.runner = TrialRunner(
|
||||
launch_web_server=True, server_port=port)
|
||||
runner = self.runner
|
||||
kwargs = {
|
||||
"stopping_criterion": {"training_iteration": 3},
|
||||
"resources": Resources(cpu=1, gpu=1),
|
||||
}
|
||||
trials = [
|
||||
Trial("__fake", **kwargs),
|
||||
Trial("__fake", **kwargs)]
|
||||
for t in trials:
|
||||
runner.add_trial(t)
|
||||
client = TuneClient("localhost:{}".format(port))
|
||||
return runner, client
|
||||
|
||||
def tearDown(self):
|
||||
print("Tearing down....")
|
||||
try:
|
||||
self.runner._server.shutdown()
|
||||
self.runner = None
|
||||
except Exception as e:
|
||||
print(e)
|
||||
ray.worker.cleanup()
|
||||
_register_all()
|
||||
|
||||
def testAddTrial(self):
|
||||
runner, client = self.basicSetup()
|
||||
for i in range(3):
|
||||
runner.step()
|
||||
spec = {
|
||||
"run": "__fake",
|
||||
"stop": {"training_iteration": 3},
|
||||
"resources": dict(cpu=1, gpu=1),
|
||||
}
|
||||
client.add_trial("test", spec)
|
||||
runner.step()
|
||||
all_trials = client.get_all_trials()["trials"]
|
||||
runner.step()
|
||||
self.assertEqual(len(all_trials), 3)
|
||||
|
||||
def testGetTrials(self):
|
||||
runner, client = self.basicSetup()
|
||||
for i in range(3):
|
||||
runner.step()
|
||||
all_trials = client.get_all_trials()["trials"]
|
||||
self.assertEqual(len(all_trials), 2)
|
||||
tid = all_trials[0]["id"]
|
||||
client.get_trial(tid)
|
||||
runner.step()
|
||||
self.assertEqual(len(all_trials), 2)
|
||||
|
||||
def testStopTrial(self):
|
||||
"""Check if Stop Trial works"""
|
||||
runner, client = self.basicSetup()
|
||||
for i in range(2):
|
||||
runner.step()
|
||||
all_trials = client.get_all_trials()["trials"]
|
||||
self.assertEqual(
|
||||
len([t for t in all_trials if t["status"] == Trial.RUNNING]), 1)
|
||||
|
||||
tid = [t for t in all_trials if t["status"] == Trial.RUNNING][0]["id"]
|
||||
client.stop_trial(tid)
|
||||
runner.step()
|
||||
|
||||
all_trials = client.get_all_trials()["trials"]
|
||||
self.assertEqual(
|
||||
len([t for t in all_trials if t["status"] == Trial.RUNNING]), 0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main(verbosity=2)
|
||||
Reference in New Issue
Block a user