mirror of
https://github.com/wassname/ray.git
synced 2026-08-12 12:20:11 +08:00
[tune] Make HyperBand Usable (#1215)
This commit is contained in:
+282
-125
@@ -3,6 +3,7 @@ from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import unittest
|
||||
import numpy as np
|
||||
|
||||
from ray.tune.hyperband import HyperBandScheduler
|
||||
from ray.tune.median_stopping_rule import MedianStoppingRule
|
||||
@@ -12,7 +13,9 @@ from ray.tune.trial_scheduler import TrialScheduler
|
||||
|
||||
|
||||
def result(t, rew):
|
||||
return TrainingResult(time_total_s=t, episode_reward_mean=rew)
|
||||
return TrainingResult(time_total_s=t,
|
||||
episode_reward_mean=rew,
|
||||
training_iteration=int(t))
|
||||
|
||||
|
||||
class EarlyStoppingSuite(unittest.TestCase):
|
||||
@@ -156,21 +159,46 @@ class HyperbandSuite(unittest.TestCase):
|
||||
"""Setup a scheduler and Runner with max Iter = 9
|
||||
|
||||
Bracketing is placed as follows:
|
||||
(3, 9);
|
||||
(5, 3) -> (2, 9);
|
||||
(9, 1) -> (3, 3) -> (1, 9); """
|
||||
sched = HyperBandScheduler(9, eta=3)
|
||||
(5, 81);
|
||||
(8, 27) -> (3, 81);
|
||||
(15, 9) -> (5, 27) -> (2, 81);
|
||||
(34, 3) -> (12, 9) -> (4, 27) -> (2, 81);
|
||||
(81, 1) -> (27, 3) -> (9, 9) -> (3, 27) -> (1, 81);"""
|
||||
sched = HyperBandScheduler()
|
||||
for i in range(num_trials):
|
||||
t = Trial("t%d" % i, "__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
runner = _MockTrialRunner()
|
||||
return sched, runner
|
||||
|
||||
def default_statistics(self):
|
||||
"""Default statistics for HyperBand"""
|
||||
sched = HyperBandScheduler()
|
||||
res = {
|
||||
str(s): {"n": sched._get_n0(s), "r": sched._get_r0(s)}
|
||||
for s in range(sched._s_max_1)
|
||||
}
|
||||
res["max_trials"] = sum(v["n"] for v in res.values())
|
||||
res["brack_count"] = sched._s_max_1
|
||||
res["s_max"] = sched._s_max_1 - 1
|
||||
return res
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
sched, _ = self.schedulerSetup(17)
|
||||
stats = self.default_statistics()
|
||||
sched, _ = self.schedulerSetup(stats["max_trials"])
|
||||
|
||||
self.assertEqual(len(sched._hyperbands), 1)
|
||||
self.assertEqual(sched._cur_band_filled(), True)
|
||||
@@ -192,7 +220,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
self.assertEqual(len(unfilled_band), 2)
|
||||
bracket = unfilled_band[-1]
|
||||
self.assertEqual(bracket.filled(), False)
|
||||
self.assertEqual(len(bracket.current_trials()), 1)
|
||||
self.assertEqual(len(bracket.current_trials()), 7)
|
||||
|
||||
return sched
|
||||
|
||||
@@ -200,19 +228,254 @@ class HyperbandSuite(unittest.TestCase):
|
||||
self.assertNotEqual(trial.status, Trial.TERMINATED)
|
||||
mock_runner._stop_trial(trial)
|
||||
|
||||
def testSuccessiveHalving(self):
|
||||
"""Setup full band, then iterate through last bracket (n=9)
|
||||
to make sure successive halving is correct."""
|
||||
def testConfigSameEta(self):
|
||||
sched = HyperBandScheduler()
|
||||
i = 0
|
||||
while not sched._cur_band_filled():
|
||||
t = Trial("t%d" % (i), "__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
i += 1
|
||||
self.assertEqual(len(sched._hyperbands[0]), 5)
|
||||
self.assertEqual(sched._hyperbands[0][0]._n, 5)
|
||||
self.assertEqual(sched._hyperbands[0][0]._r, 81)
|
||||
self.assertEqual(sched._hyperbands[0][-1]._n, 81)
|
||||
self.assertEqual(sched._hyperbands[0][-1]._r, 1)
|
||||
|
||||
sched, mock_runner = self.schedulerSetup(17)
|
||||
filled_band = sched._hyperbands[0][-1]
|
||||
big_bracket = filled_band
|
||||
sched = HyperBandScheduler(max_t=810)
|
||||
i = 0
|
||||
while not sched._cur_band_filled():
|
||||
t = Trial("t%d" % (i), "__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
i += 1
|
||||
self.assertEqual(len(sched._hyperbands[0]), 5)
|
||||
self.assertEqual(sched._hyperbands[0][0]._n, 5)
|
||||
self.assertEqual(sched._hyperbands[0][0]._r, 810)
|
||||
self.assertEqual(sched._hyperbands[0][-1]._n, 81)
|
||||
self.assertEqual(sched._hyperbands[0][-1]._r, 10)
|
||||
|
||||
def testConfigSameEtaSmall(self):
|
||||
sched = HyperBandScheduler(max_t=1)
|
||||
i = 0
|
||||
while len(sched._hyperbands) < 2:
|
||||
t = Trial("t%d" % (i), "__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
i += 1
|
||||
self.assertEqual(len(sched._hyperbands[0]), 5)
|
||||
self.assertTrue(all(v is None for v in sched._hyperbands[0][1:]))
|
||||
|
||||
def testSuccessiveHalving(self):
|
||||
"""Setup full band, then iterate through last bracket (n=81)
|
||||
to make sure successive halving is correct."""
|
||||
stats = self.default_statistics()
|
||||
sched, mock_runner = self.schedulerSetup(stats["max_trials"])
|
||||
big_bracket = sched._state["bracket"]
|
||||
cur_units = stats[str(stats["s_max"])]["r"]
|
||||
# The last bracket will downscale 4 times
|
||||
for x in range(stats["brack_count"] - 1):
|
||||
trials = big_bracket.current_trials()
|
||||
current_length = len(trials)
|
||||
for trl in trials:
|
||||
mock_runner._launch_trial(trl)
|
||||
|
||||
# Provides results from 0 to 8 in order, keeping last one running
|
||||
for i, trl in enumerate(trials):
|
||||
action = sched.on_trial_result(
|
||||
mock_runner, trl, result(cur_units, i))
|
||||
if i < current_length - 1:
|
||||
self.assertEqual(action, TrialScheduler.PAUSE)
|
||||
self.process(trl, mock_runner, action)
|
||||
|
||||
self.assertEqual(action, TrialScheduler.CONTINUE)
|
||||
new_length = len(big_bracket.current_trials())
|
||||
self.assertEqual(new_length, self.downscale(current_length, sched))
|
||||
cur_units += int(cur_units * sched._eta)
|
||||
self.assertEqual(len(big_bracket.current_trials()), 1)
|
||||
|
||||
def testHalvingStop(self):
|
||||
stats = self.default_statistics()
|
||||
num_trials = stats[str(0)]["n"] + stats[str(1)]["n"]
|
||||
sched, mock_runner = self.schedulerSetup(num_trials)
|
||||
big_bracket = sched._state["bracket"]
|
||||
for trl in big_bracket.current_trials():
|
||||
mock_runner._launch_trial(trl)
|
||||
|
||||
# # Provides result in reverse order, killing the last one
|
||||
cur_units = stats[str(1)]["r"]
|
||||
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)
|
||||
|
||||
self.assertEqual(action, TrialScheduler.STOP)
|
||||
|
||||
def testContinueLastOne(self):
|
||||
stats = self.default_statistics()
|
||||
num_trials = stats[str(0)]["n"]
|
||||
sched, mock_runner = self.schedulerSetup(num_trials)
|
||||
big_bracket = sched._state["bracket"]
|
||||
for trl in big_bracket.current_trials():
|
||||
mock_runner._launch_trial(trl)
|
||||
|
||||
# # Provides result in reverse order, killing the last one
|
||||
cur_units = stats[str(0)]["r"]
|
||||
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)
|
||||
|
||||
self.assertEqual(action, TrialScheduler.CONTINUE)
|
||||
|
||||
for x in range(100):
|
||||
action = sched.on_trial_result(
|
||||
mock_runner, trl, result(cur_units + x, 10))
|
||||
self.assertEqual(action, TrialScheduler.CONTINUE)
|
||||
|
||||
def testTrialErrored(self):
|
||||
"""If a trial errored, make sure successive halving still happens"""
|
||||
stats = self.default_statistics()
|
||||
trial_count = stats[str(0)]["n"] + 3
|
||||
sched, mock_runner = self.schedulerSetup(trial_count)
|
||||
t1, t2, t3 = sched._state["bracket"].current_trials()
|
||||
for t in [t1, t2, t3]:
|
||||
mock_runner._launch_trial(t)
|
||||
|
||||
sched.on_trial_error(mock_runner, t3)
|
||||
self.assertEqual(
|
||||
TrialScheduler.PAUSE,
|
||||
sched.on_trial_result(
|
||||
mock_runner, t1, result(stats[str(1)]["r"], 10)))
|
||||
self.assertEqual(
|
||||
TrialScheduler.CONTINUE,
|
||||
sched.on_trial_result(
|
||||
mock_runner, t2, result(stats[str(1)]["r"], 10)))
|
||||
|
||||
def testTrialErrored2(self):
|
||||
"""Check successive halving happened even when last trial failed"""
|
||||
stats = self.default_statistics()
|
||||
trial_count = stats[str(0)]["n"] + stats[str(1)]["n"]
|
||||
sched, mock_runner = self.schedulerSetup(trial_count)
|
||||
trials = sched._state["bracket"].current_trials()
|
||||
for t in trials[:-1]:
|
||||
mock_runner._launch_trial(t)
|
||||
sched.on_trial_result(
|
||||
mock_runner, t, result(stats[str(1)]["r"], 10))
|
||||
|
||||
mock_runner._launch_trial(trials[-1])
|
||||
sched.on_trial_error(mock_runner, trials[-1])
|
||||
self.assertEqual(len(sched._state["bracket"].current_trials()),
|
||||
self.downscale(stats[str(1)]["n"], sched))
|
||||
|
||||
def testTrialEndedEarly(self):
|
||||
"""Check successive halving happened even when one trial failed"""
|
||||
stats = self.default_statistics()
|
||||
trial_count = stats[str(0)]["n"] + 3
|
||||
sched, mock_runner = self.schedulerSetup(trial_count)
|
||||
|
||||
t1, t2, t3 = sched._state["bracket"].current_trials()
|
||||
for t in [t1, t2, t3]:
|
||||
mock_runner._launch_trial(t)
|
||||
|
||||
sched.on_trial_complete(mock_runner, t3, result(1, 12))
|
||||
self.assertEqual(
|
||||
TrialScheduler.PAUSE,
|
||||
sched.on_trial_result(
|
||||
mock_runner, t1, result(stats[str(1)]["r"], 10)))
|
||||
self.assertEqual(
|
||||
TrialScheduler.CONTINUE,
|
||||
sched.on_trial_result(
|
||||
mock_runner, t2, result(stats[str(1)]["r"], 10)))
|
||||
|
||||
def testTrialEndedEarly2(self):
|
||||
"""Check successive halving happened even when last trial failed"""
|
||||
stats = self.default_statistics()
|
||||
trial_count = stats[str(0)]["n"] + stats[str(1)]["n"]
|
||||
sched, mock_runner = self.schedulerSetup(trial_count)
|
||||
trials = sched._state["bracket"].current_trials()
|
||||
for t in trials[:-1]:
|
||||
mock_runner._launch_trial(t)
|
||||
sched.on_trial_result(
|
||||
mock_runner, t, result(stats[str(1)]["r"], 10))
|
||||
|
||||
mock_runner._launch_trial(trials[-1])
|
||||
sched.on_trial_complete(mock_runner, trials[-1], result(100, 12))
|
||||
self.assertEqual(len(sched._state["bracket"].current_trials()),
|
||||
self.downscale(stats[str(1)]["n"], sched))
|
||||
|
||||
def testAddAfterHalving(self):
|
||||
stats = self.default_statistics()
|
||||
trial_count = stats[str(0)]["n"] + 1
|
||||
sched, mock_runner = self.schedulerSetup(trial_count)
|
||||
bracket_trials = sched._state["bracket"].current_trials()
|
||||
init_units = stats[str(1)]["r"]
|
||||
|
||||
for t in bracket_trials:
|
||||
mock_runner._launch_trial(t)
|
||||
|
||||
for i, t in enumerate(bracket_trials):
|
||||
status = sched.on_trial_result(
|
||||
mock_runner, t, result(init_units, i))
|
||||
self.assertEqual(status, TrialScheduler.CONTINUE)
|
||||
t = Trial("t%d" % 100, "__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
mock_runner._launch_trial(t)
|
||||
self.assertEqual(len(sched._state["bracket"].current_trials()), 2)
|
||||
|
||||
# Make sure that newly added trial gets fair computation (not just 1)
|
||||
self.assertEqual(
|
||||
TrialScheduler.CONTINUE,
|
||||
sched.on_trial_result(mock_runner, t, result(init_units, 12)))
|
||||
new_units = init_units + int(init_units * sched._eta)
|
||||
self.assertEqual(
|
||||
TrialScheduler.PAUSE,
|
||||
sched.on_trial_result(mock_runner, t, result(new_units, 12)))
|
||||
|
||||
def testAlternateMetrics(self):
|
||||
"""Checking that alternate metrics will pass."""
|
||||
|
||||
def result2(t, rew):
|
||||
return TrainingResult(time_total_s=t, neg_mean_loss=rew)
|
||||
|
||||
sched = HyperBandScheduler(
|
||||
time_attr='time_total_s', reward_attr='neg_mean_loss')
|
||||
stats = self.default_statistics()
|
||||
|
||||
for i in range(stats["max_trials"]):
|
||||
t = Trial("t%d" % i, "__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
runner = _MockTrialRunner()
|
||||
|
||||
big_bracket = sched._hyperbands[0][-1]
|
||||
|
||||
for trl in big_bracket.current_trials():
|
||||
runner._launch_trial(trl)
|
||||
current_length = len(big_bracket.current_trials())
|
||||
|
||||
# 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)
|
||||
|
||||
new_length = len(big_bracket.current_trials())
|
||||
self.assertEqual(status, TrialScheduler.CONTINUE)
|
||||
self.assertEqual(new_length, self.downscale(current_length, sched))
|
||||
|
||||
def testJumpingTime(self):
|
||||
sched, mock_runner = self.schedulerSetup(81)
|
||||
big_bracket = sched._hyperbands[0][-1]
|
||||
|
||||
for trl in big_bracket.current_trials():
|
||||
mock_runner._launch_trial(trl)
|
||||
|
||||
# Provides results from 0 to 8 in order, keeping the last one running
|
||||
for i, trl in enumerate(big_bracket.current_trials()):
|
||||
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
|
||||
@@ -222,117 +485,11 @@ class HyperbandSuite(unittest.TestCase):
|
||||
self.assertNotEqual(trl.status, Trial.TERMINATED)
|
||||
self.stopTrial(trl, mock_runner)
|
||||
|
||||
status = sched.on_trial_result(mock_runner, jump, result(4, i))
|
||||
self.assertEqual(status, TrialScheduler.PAUSE)
|
||||
|
||||
current_length = len(big_bracket.current_trials())
|
||||
self.assertEqual(status, TrialScheduler.CONTINUE)
|
||||
self.assertEqual(current_length, 3)
|
||||
|
||||
# Techincally only need to launch 2/3, as one is already running
|
||||
for trl in big_bracket.current_trials():
|
||||
mock_runner._launch_trial(trl)
|
||||
|
||||
# Provides results from 2 to 0 in order, killing the last one
|
||||
for i, trl in reversed(list(enumerate(big_bracket.current_trials()))):
|
||||
for j in range(3):
|
||||
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.stopTrial(trl, mock_runner)
|
||||
|
||||
self.assertEqual(status, TrialScheduler.STOP)
|
||||
trl = big_bracket.current_trials()[0]
|
||||
for i in range(9):
|
||||
status = sched.on_trial_result(mock_runner, trl, result(1, i))
|
||||
self.assertEqual(status, TrialScheduler.STOP)
|
||||
self.assertEqual(len(big_bracket.current_trials()), 0)
|
||||
self.assertEqual(sched._num_stopped, 9)
|
||||
|
||||
def testScheduling(self):
|
||||
"""Setup two bands, then make sure all trials are running"""
|
||||
sched = self.advancedSetup()
|
||||
mock_runner = _MockTrialRunner()
|
||||
trl = sched.choose_trial_to_run(mock_runner)
|
||||
while trl:
|
||||
# If band iteration > 0, make sure first band is all running
|
||||
if sched._trial_info[trl][1] > 0:
|
||||
first_band = sched._hyperbands[0]
|
||||
trials = [t for b in first_band for t in b._live_trials]
|
||||
self.assertEqual(
|
||||
all(t.status == Trial.RUNNING for t in trials),
|
||||
True)
|
||||
mock_runner._launch_trial(trl)
|
||||
res = sched.on_trial_result(mock_runner, trl, result(1, 10))
|
||||
if res is TrialScheduler.PAUSE:
|
||||
mock_runner._pause_trial(trl)
|
||||
trl = sched.choose_trial_to_run(mock_runner)
|
||||
|
||||
self.assertEqual(
|
||||
all(t.status == Trial.RUNNING for t in trials), True)
|
||||
|
||||
def testTrialErrored(self):
|
||||
sched, mock_runner = self.schedulerSetup(10)
|
||||
t1, t2 = sched._state["bracket"].current_trials()
|
||||
mock_runner._launch_trial(t1)
|
||||
mock_runner._launch_trial(t2)
|
||||
|
||||
sched.on_trial_error(mock_runner, t2)
|
||||
self.assertEqual(
|
||||
TrialScheduler.CONTINUE,
|
||||
sched.on_trial_result(mock_runner, t1, result(1, 10)))
|
||||
|
||||
def testTrialErrored2(self):
|
||||
"""Check successive halving happened even when last trial failed"""
|
||||
sched, mock_runner = self.schedulerSetup(17)
|
||||
trials = sched._state["bracket"].current_trials()
|
||||
self.assertEqual(len(trials), 9)
|
||||
for t in trials[:-1]:
|
||||
mock_runner._launch_trial(t)
|
||||
sched.on_trial_result(mock_runner, t, result(1, 10))
|
||||
|
||||
mock_runner._launch_trial(trials[-1])
|
||||
sched.on_trial_error(mock_runner, trials[-1])
|
||||
self.assertEqual(len(sched._state["bracket"].current_trials()), 3)
|
||||
|
||||
def testTrialEndedEarly(self):
|
||||
sched, mock_runner = self.schedulerSetup(10)
|
||||
trials = sched._state["bracket"].current_trials()
|
||||
for t in trials:
|
||||
mock_runner._launch_trial(t)
|
||||
|
||||
sched.on_trial_complete(mock_runner, trials[-1], result(1, 12))
|
||||
self.assertEqual(
|
||||
TrialScheduler.CONTINUE,
|
||||
sched.on_trial_result(mock_runner, trials[0], result(1, 12)))
|
||||
|
||||
def testTrialEndedEarly2(self):
|
||||
"""Check successive halving happened even when last trial finished"""
|
||||
sched, mock_runner = self.schedulerSetup(17)
|
||||
trials = sched._state["bracket"].current_trials()
|
||||
self.assertEqual(len(trials), 9)
|
||||
for t in trials[:-1]:
|
||||
mock_runner._launch_trial(t)
|
||||
sched.on_trial_result(mock_runner, t, result(1, 10))
|
||||
|
||||
mock_runner._launch_trial(trials[-1])
|
||||
sched.on_trial_complete(mock_runner, trials[-1], result(1, 12))
|
||||
self.assertEqual(len(sched._state["bracket"].current_trials()), 3)
|
||||
|
||||
def testAddAfterHalving(self):
|
||||
sched, mock_runner = self.schedulerSetup(10)
|
||||
bracket_trials = sched._state["bracket"].current_trials()
|
||||
|
||||
for t in bracket_trials:
|
||||
mock_runner._launch_trial(t)
|
||||
|
||||
for i, t in enumerate(bracket_trials):
|
||||
res = sched.on_trial_result(
|
||||
mock_runner, t, result(1, i))
|
||||
self.assertEqual(res, TrialScheduler.CONTINUE)
|
||||
t = Trial("t%d" % 5, "__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
self.assertEqual(3 + 1, sched._state["bracket"]._live_trials[t][1])
|
||||
self.assertLess(current_length, 27)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user