mirror of
https://github.com/wassname/ray.git
synced 2026-08-12 12:20:11 +08:00
[tune] Support user-defined trainable functions / classes / envs with a shared object registry (#1226)
This commit is contained in:
@@ -59,83 +59,83 @@ python $ROOT_DIR/multi_node_docker_test.py \
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env PongDeterministic-v0 \
|
||||
--alg A3C \
|
||||
--run A3C \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"num_workers": 16}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env CartPole-v1 \
|
||||
--alg PPO \
|
||||
--run PPO \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"kl_coeff": 1.0, "num_sgd_iter": 10, "sgd_stepsize": 1e-4, "sgd_batchsize": 64, "timesteps_per_batch": 2000, "num_workers": 1, "model": {"free_log_std": true}}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env CartPole-v1 \
|
||||
--alg PPO \
|
||||
--run PPO \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"kl_coeff": 1.0, "num_sgd_iter": 10, "sgd_stepsize": 1e-4, "sgd_batchsize": 64, "timesteps_per_batch": 2000, "num_workers": 1, "use_gae": false}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env Pendulum-v0 \
|
||||
--alg ES \
|
||||
--run ES \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"stepsize": 0.01, "episodes_per_batch": 20, "timesteps_per_batch": 100}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env Pong-v0 \
|
||||
--alg ES \
|
||||
--run ES \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"stepsize": 0.01, "episodes_per_batch": 20, "timesteps_per_batch": 100}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env CartPole-v0 \
|
||||
--alg A3C \
|
||||
--run A3C \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"use_lstm": false}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env CartPole-v0 \
|
||||
--alg DQN \
|
||||
--run DQN \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"lr": 1e-3, "schedule_max_timesteps": 100000, "exploration_fraction": 0.1, "exploration_final_eps": 0.02, "dueling": false, "hiddens": [], "model": {"fcnet_hiddens": [64], "fcnet_activation": "relu"}}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env FrozenLake-v0 \
|
||||
--alg DQN \
|
||||
--run DQN \
|
||||
--stop '{"training_iteration": 2}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env FrozenLake-v0 \
|
||||
--alg PPO \
|
||||
--run PPO \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"num_sgd_iter": 10, "sgd_batchsize": 64, "timesteps_per_batch": 1000, "num_workers": 1}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env PongDeterministic-v4 \
|
||||
--alg DQN \
|
||||
--run DQN \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"lr": 1e-4, "schedule_max_timesteps": 2000000, "buffer_size": 10000, "exploration_fraction": 0.1, "exploration_final_eps": 0.01, "sample_batch_size": 4, "learning_starts": 10000, "target_network_update_freq": 1000, "gamma": 0.99, "prioritized_replay": true}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env MontezumaRevenge-v0 \
|
||||
--alg PPO \
|
||||
--run PPO \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"kl_coeff": 1.0, "num_sgd_iter": 10, "sgd_stepsize": 1e-4, "sgd_batchsize": 64, "timesteps_per_batch": 2000, "num_workers": 1, "model": {"dim": 40, "conv_filters": [[16, [8, 8], 4], [32, [4, 4], 2], [512, [5, 5], 1]]}, "extra_frameskip": 4}'
|
||||
|
||||
docker run --shm-size=10G --memory=10G $DOCKER_SHA \
|
||||
python /ray/python/ray/rllib/train.py \
|
||||
--env PongDeterministic-v4 \
|
||||
--alg A3C \
|
||||
--run A3C \
|
||||
--stop '{"training_iteration": 2}' \
|
||||
--config '{"num_workers": 2, "use_lstm": false, "use_pytorch": true, "model": {"grayscale": true, "zero_mean": false, "dim": 80, "channel_major": true}}'
|
||||
|
||||
|
||||
+196
-30
@@ -2,38 +2,201 @@ from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import unittest
|
||||
import os
|
||||
import time
|
||||
import unittest
|
||||
|
||||
import ray
|
||||
from ray.rllib import _register_all
|
||||
|
||||
from ray.tune import Trainable, TuneError
|
||||
from ray.tune import register_env, register_trainable, run_experiments
|
||||
from ray.tune.registry import _default_registry, TRAINABLE_CLASS
|
||||
from ray.tune.trial import Trial, Resources
|
||||
from ray.tune.trial_runner import TrialRunner
|
||||
from ray.tune.variant_generator import generate_trials, grid_search, \
|
||||
RecursiveDependencyError
|
||||
|
||||
|
||||
class TrainableFunctionApiTest(unittest.TestCase):
|
||||
def tearDown(self):
|
||||
ray.worker.cleanup()
|
||||
_register_all() # re-register the evicted objects
|
||||
|
||||
def testRegisterEnv(self):
|
||||
register_env("foo", lambda: None)
|
||||
self.assertRaises(TypeError, lambda: register_env("foo", 2))
|
||||
|
||||
def testRegisterTrainable(self):
|
||||
def train(config, reporter):
|
||||
pass
|
||||
|
||||
class A(object):
|
||||
pass
|
||||
|
||||
class B(Trainable):
|
||||
pass
|
||||
|
||||
register_trainable("foo", train)
|
||||
register_trainable("foo", B)
|
||||
self.assertRaises(TypeError, lambda: register_trainable("foo", B()))
|
||||
self.assertRaises(TypeError, lambda: register_trainable("foo", A))
|
||||
|
||||
def testRewriteEnv(self):
|
||||
def train(config, reporter):
|
||||
reporter(timesteps_total=1)
|
||||
register_trainable("f1", train)
|
||||
|
||||
[trial] = run_experiments({"foo": {
|
||||
"run": "f1",
|
||||
"env": "CartPole-v0",
|
||||
}})
|
||||
self.assertEqual(trial.config["env"], "CartPole-v0")
|
||||
|
||||
def testConfigPurity(self):
|
||||
def train(config, reporter):
|
||||
assert config == {"a": "b"}, config
|
||||
reporter(timesteps_total=1)
|
||||
register_trainable("f1", train)
|
||||
run_experiments({"foo": {
|
||||
"run": "f1",
|
||||
"config": {"a": "b"},
|
||||
}})
|
||||
|
||||
def testBadParams(self):
|
||||
def f():
|
||||
run_experiments({"foo": {}})
|
||||
self.assertRaises(TuneError, f)
|
||||
|
||||
def testBadParams2(self):
|
||||
def f():
|
||||
run_experiments({"foo": {
|
||||
"bah": "this param is not allowed",
|
||||
}})
|
||||
self.assertRaises(TuneError, f)
|
||||
|
||||
def testBadParams3(self):
|
||||
def f():
|
||||
run_experiments({"foo": {
|
||||
"run": grid_search("invalid grid search"),
|
||||
}})
|
||||
self.assertRaises(TuneError, f)
|
||||
|
||||
def testBadParams4(self):
|
||||
def f():
|
||||
run_experiments({"foo": {
|
||||
"run": "asdf",
|
||||
}})
|
||||
self.assertRaises(TuneError, f)
|
||||
|
||||
def testBadParams5(self):
|
||||
def f():
|
||||
run_experiments({"foo": {
|
||||
"run": "PPO",
|
||||
"stop": {"asdf": 1}
|
||||
}})
|
||||
self.assertRaises(TuneError, f)
|
||||
|
||||
def testBadParams6(self):
|
||||
def f():
|
||||
run_experiments({"foo": {
|
||||
"run": "PPO",
|
||||
"resources": {"asdf": 1}
|
||||
}})
|
||||
self.assertRaises(TuneError, f)
|
||||
|
||||
def testBadReturn(self):
|
||||
def train(config, reporter):
|
||||
reporter()
|
||||
register_trainable("f1", train)
|
||||
|
||||
def f():
|
||||
run_experiments({"foo": {
|
||||
"run": "f1",
|
||||
"config": {
|
||||
"script_min_iter_time_s": 0,
|
||||
},
|
||||
}})
|
||||
self.assertRaises(TuneError, f)
|
||||
|
||||
def testEarlyReturn(self):
|
||||
def train(config, reporter):
|
||||
reporter(timesteps_total=100, done=True)
|
||||
time.sleep(99999)
|
||||
register_trainable("f1", train)
|
||||
[trial] = run_experiments({"foo": {
|
||||
"run": "f1",
|
||||
"config": {
|
||||
"script_min_iter_time_s": 0,
|
||||
},
|
||||
}})
|
||||
self.assertEqual(trial.status, Trial.TERMINATED)
|
||||
self.assertEqual(trial.last_result.timesteps_total, 100)
|
||||
|
||||
def testAbruptReturn(self):
|
||||
def train(config, reporter):
|
||||
reporter(timesteps_total=100)
|
||||
register_trainable("f1", train)
|
||||
[trial] = run_experiments({"foo": {
|
||||
"run": "f1",
|
||||
"config": {
|
||||
"script_min_iter_time_s": 0,
|
||||
},
|
||||
}})
|
||||
self.assertEqual(trial.status, Trial.TERMINATED)
|
||||
self.assertEqual(trial.last_result.timesteps_total, 100)
|
||||
|
||||
def testErrorReturn(self):
|
||||
def train(config, reporter):
|
||||
raise Exception("uh oh")
|
||||
register_trainable("f1", train)
|
||||
|
||||
def f():
|
||||
run_experiments({"foo": {
|
||||
"run": "f1",
|
||||
"config": {
|
||||
"script_min_iter_time_s": 0,
|
||||
},
|
||||
}})
|
||||
self.assertRaises(TuneError, f)
|
||||
|
||||
def testSuccess(self):
|
||||
def train(config, reporter):
|
||||
for i in range(100):
|
||||
reporter(timesteps_total=i)
|
||||
register_trainable("f1", train)
|
||||
[trial] = run_experiments({"foo": {
|
||||
"run": "f1",
|
||||
"config": {
|
||||
"script_min_iter_time_s": 0,
|
||||
},
|
||||
}})
|
||||
self.assertEqual(trial.status, Trial.TERMINATED)
|
||||
self.assertEqual(trial.last_result.timesteps_total, 99)
|
||||
|
||||
|
||||
class VariantGeneratorTest(unittest.TestCase):
|
||||
def testParseToTrials(self):
|
||||
trials = generate_trials({
|
||||
"env": "Pong-v0",
|
||||
"alg": "PPO",
|
||||
"run": "PPO",
|
||||
"repeat": 2,
|
||||
"config": {
|
||||
"env": "Pong-v0",
|
||||
"foo": "bar"
|
||||
},
|
||||
}, "tune-pong")
|
||||
trials = list(trials)
|
||||
self.assertEqual(len(trials), 2)
|
||||
self.assertEqual(trials[0].env_name, "Pong-v0")
|
||||
self.assertEqual(trials[0].config, {"foo": "bar"})
|
||||
self.assertEqual(trials[0].alg, "PPO")
|
||||
self.assertEqual(str(trials[0]), "PPO_Pong-v0_0")
|
||||
self.assertEqual(trials[0].config, {"foo": "bar", "env": "Pong-v0"})
|
||||
self.assertEqual(trials[0].trainable_name, "PPO")
|
||||
self.assertEqual(trials[0].experiment_tag, "0")
|
||||
self.assertEqual(trials[0].local_dir, "/tmp/ray/tune-pong")
|
||||
self.assertEqual(trials[1].experiment_tag, "1")
|
||||
|
||||
def testEval(self):
|
||||
trials = generate_trials({
|
||||
"env": "Pong-v0",
|
||||
"run": "PPO",
|
||||
"config": {
|
||||
"foo": {
|
||||
"eval": "2 + 2"
|
||||
@@ -48,7 +211,7 @@ class VariantGeneratorTest(unittest.TestCase):
|
||||
|
||||
def testGridSearch(self):
|
||||
trials = generate_trials({
|
||||
"env": "Pong-v0",
|
||||
"run": "PPO",
|
||||
"config": {
|
||||
"bar": {
|
||||
"grid_search": [True, False]
|
||||
@@ -71,7 +234,7 @@ class VariantGeneratorTest(unittest.TestCase):
|
||||
|
||||
def testGridSearchAndEval(self):
|
||||
trials = generate_trials({
|
||||
"env": "Pong-v0",
|
||||
"run": "PPO",
|
||||
"config": {
|
||||
"qux": lambda spec: 2 + 2,
|
||||
"bar": grid_search([True, False]),
|
||||
@@ -85,7 +248,7 @@ class VariantGeneratorTest(unittest.TestCase):
|
||||
|
||||
def testConditionResolution(self):
|
||||
trials = generate_trials({
|
||||
"env": "Pong-v0",
|
||||
"run": "PPO",
|
||||
"config": {
|
||||
"x": 1,
|
||||
"y": lambda spec: spec.config.x + 1,
|
||||
@@ -98,7 +261,7 @@ class VariantGeneratorTest(unittest.TestCase):
|
||||
|
||||
def testDependentLambda(self):
|
||||
trials = generate_trials({
|
||||
"env": "Pong-v0",
|
||||
"run": "PPO",
|
||||
"config": {
|
||||
"x": grid_search([1, 2]),
|
||||
"y": lambda spec: spec.config.x * 100,
|
||||
@@ -111,7 +274,7 @@ class VariantGeneratorTest(unittest.TestCase):
|
||||
|
||||
def testDependentGridSearch(self):
|
||||
trials = generate_trials({
|
||||
"env": "Pong-v0",
|
||||
"run": "PPO",
|
||||
"config": {
|
||||
"x": grid_search([
|
||||
lambda spec: spec.config.y * 100,
|
||||
@@ -128,7 +291,7 @@ class VariantGeneratorTest(unittest.TestCase):
|
||||
def testRecursiveDep(self):
|
||||
try:
|
||||
list(generate_trials({
|
||||
"env": "Pong-v0",
|
||||
"run": "PPO",
|
||||
"config": {
|
||||
"foo": lambda spec: spec.config.foo,
|
||||
},
|
||||
@@ -142,10 +305,11 @@ class VariantGeneratorTest(unittest.TestCase):
|
||||
class TrialRunnerTest(unittest.TestCase):
|
||||
def tearDown(self):
|
||||
ray.worker.cleanup()
|
||||
_register_all() # re-register the evicted objects
|
||||
|
||||
def testTrialStatus(self):
|
||||
ray.init()
|
||||
trial = Trial("CartPole-v0", "__fake")
|
||||
trial = Trial("__fake")
|
||||
self.assertEqual(trial.status, Trial.PENDING)
|
||||
trial.start()
|
||||
self.assertEqual(trial.status, Trial.RUNNING)
|
||||
@@ -156,11 +320,12 @@ class TrialRunnerTest(unittest.TestCase):
|
||||
|
||||
def testTrialErrorOnStart(self):
|
||||
ray.init()
|
||||
trial = Trial("CartPole-v0", "asdf")
|
||||
_default_registry.register(TRAINABLE_CLASS, "asdf", None)
|
||||
trial = Trial("asdf")
|
||||
try:
|
||||
trial.start()
|
||||
except Exception as e:
|
||||
self.assertIn("Unknown algorithm", str(e))
|
||||
self.assertIn("a class", str(e))
|
||||
|
||||
def testResourceScheduler(self):
|
||||
ray.init(num_cpus=4, num_gpus=1)
|
||||
@@ -170,8 +335,8 @@ class TrialRunnerTest(unittest.TestCase):
|
||||
"resources": Resources(cpu=1, gpu=1),
|
||||
}
|
||||
trials = [
|
||||
Trial("CartPole-v0", "__fake", **kwargs),
|
||||
Trial("CartPole-v0", "__fake", **kwargs)]
|
||||
Trial("__fake", **kwargs),
|
||||
Trial("__fake", **kwargs)]
|
||||
for t in trials:
|
||||
runner.add_trial(t)
|
||||
|
||||
@@ -199,8 +364,8 @@ class TrialRunnerTest(unittest.TestCase):
|
||||
"resources": Resources(cpu=1, gpu=1),
|
||||
}
|
||||
trials = [
|
||||
Trial("CartPole-v0", "__fake", **kwargs),
|
||||
Trial("CartPole-v0", "__fake", **kwargs)]
|
||||
Trial("__fake", **kwargs),
|
||||
Trial("__fake", **kwargs)]
|
||||
for t in trials:
|
||||
runner.add_trial(t)
|
||||
|
||||
@@ -227,9 +392,10 @@ class TrialRunnerTest(unittest.TestCase):
|
||||
"stopping_criterion": {"training_iteration": 1},
|
||||
"resources": Resources(cpu=1, gpu=1),
|
||||
}
|
||||
_default_registry.register(TRAINABLE_CLASS, "asdf", None)
|
||||
trials = [
|
||||
Trial("CartPole-v0", "asdf", **kwargs),
|
||||
Trial("CartPole-v0", "__fake", **kwargs)]
|
||||
Trial("asdf", **kwargs),
|
||||
Trial("__fake", **kwargs)]
|
||||
for t in trials:
|
||||
runner.add_trial(t)
|
||||
|
||||
@@ -248,17 +414,17 @@ class TrialRunnerTest(unittest.TestCase):
|
||||
"stopping_criterion": {"training_iteration": 1},
|
||||
"resources": Resources(cpu=1, gpu=1),
|
||||
}
|
||||
runner.add_trial(Trial("CartPole-v0", "__fake", **kwargs))
|
||||
runner.add_trial(Trial("__fake", **kwargs))
|
||||
trials = runner.get_trials()
|
||||
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.RUNNING)
|
||||
self.assertEqual(ray.get(trials[0].agent.set_info.remote(1)), 1)
|
||||
self.assertEqual(ray.get(trials[0].runner.set_info.remote(1)), 1)
|
||||
|
||||
path = trials[0].checkpoint()
|
||||
kwargs["restore_path"] = path
|
||||
|
||||
runner.add_trial(Trial("CartPole-v0", "__fake", **kwargs))
|
||||
runner.add_trial(Trial("__fake", **kwargs))
|
||||
trials = runner.get_trials()
|
||||
|
||||
runner.step()
|
||||
@@ -268,7 +434,7 @@ class TrialRunnerTest(unittest.TestCase):
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.TERMINATED)
|
||||
self.assertEqual(trials[1].status, Trial.RUNNING)
|
||||
self.assertEqual(ray.get(trials[1].agent.get_info.remote()), 1)
|
||||
self.assertEqual(ray.get(trials[1].runner.get_info.remote()), 1)
|
||||
self.addCleanup(os.remove, path)
|
||||
|
||||
def testPauseThenResume(self):
|
||||
@@ -278,14 +444,14 @@ class TrialRunnerTest(unittest.TestCase):
|
||||
"stopping_criterion": {"training_iteration": 2},
|
||||
"resources": Resources(cpu=1, gpu=1),
|
||||
}
|
||||
runner.add_trial(Trial("CartPole-v0", "__fake", **kwargs))
|
||||
runner.add_trial(Trial("__fake", **kwargs))
|
||||
trials = runner.get_trials()
|
||||
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.RUNNING)
|
||||
self.assertEqual(ray.get(trials[0].agent.get_info.remote()), None)
|
||||
self.assertEqual(ray.get(trials[0].runner.get_info.remote()), None)
|
||||
|
||||
self.assertEqual(ray.get(trials[0].agent.set_info.remote(1)), 1)
|
||||
self.assertEqual(ray.get(trials[0].runner.set_info.remote(1)), 1)
|
||||
|
||||
trials[0].pause()
|
||||
self.assertEqual(trials[0].status, Trial.PAUSED)
|
||||
@@ -295,7 +461,7 @@ class TrialRunnerTest(unittest.TestCase):
|
||||
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.RUNNING)
|
||||
self.assertEqual(ray.get(trials[0].agent.get_info.remote()), 1)
|
||||
self.assertEqual(ray.get(trials[0].runner.get_info.remote()), 1)
|
||||
|
||||
runner.step()
|
||||
self.assertEqual(trials[0].status, Trial.TERMINATED)
|
||||
|
||||
@@ -20,8 +20,8 @@ def result(t, rew):
|
||||
|
||||
class EarlyStoppingSuite(unittest.TestCase):
|
||||
def basicSetup(self, rule):
|
||||
t1 = Trial("t1", "PPO") # mean is 450, max 900, t_max=10
|
||||
t2 = Trial("t2", "PPO") # mean is 450, max 450, t_max=5
|
||||
t1 = Trial("PPO") # mean is 450, max 900, t_max=10
|
||||
t2 = Trial("PPO") # mean is 450, max 450, t_max=5
|
||||
for i in range(10):
|
||||
self.assertEqual(
|
||||
rule.on_trial_result(None, t1, result(i, i * 100)),
|
||||
@@ -62,7 +62,7 @@ class EarlyStoppingSuite(unittest.TestCase):
|
||||
t1, t2 = self.basicSetup(rule)
|
||||
rule.on_trial_complete(None, t1, result(10, 1000))
|
||||
rule.on_trial_complete(None, t2, result(10, 1000))
|
||||
t3 = Trial("t3", "PPO")
|
||||
t3 = Trial("PPO")
|
||||
self.assertEqual(
|
||||
rule.on_trial_result(None, t3, result(1, 10)),
|
||||
TrialScheduler.CONTINUE)
|
||||
@@ -77,7 +77,7 @@ class EarlyStoppingSuite(unittest.TestCase):
|
||||
rule = MedianStoppingRule(grace_period=0, min_samples_required=2)
|
||||
t1, t2 = self.basicSetup(rule)
|
||||
rule.on_trial_complete(None, t1, result(10, 1000))
|
||||
t3 = Trial("t3", "PPO")
|
||||
t3 = Trial("PPO")
|
||||
self.assertEqual(
|
||||
rule.on_trial_result(None, t3, result(3, 10)),
|
||||
TrialScheduler.CONTINUE)
|
||||
@@ -91,7 +91,7 @@ class EarlyStoppingSuite(unittest.TestCase):
|
||||
t1, t2 = self.basicSetup(rule)
|
||||
rule.on_trial_complete(None, t1, result(10, 1000))
|
||||
rule.on_trial_complete(None, t2, result(10, 1000))
|
||||
t3 = Trial("t3", "PPO")
|
||||
t3 = Trial("PPO")
|
||||
self.assertEqual(
|
||||
rule.on_trial_result(None, t3, result(1, 260)),
|
||||
TrialScheduler.CONTINUE)
|
||||
@@ -105,7 +105,7 @@ class EarlyStoppingSuite(unittest.TestCase):
|
||||
t1, t2 = self.basicSetup(rule)
|
||||
rule.on_trial_complete(None, t1, result(10, 1000))
|
||||
rule.on_trial_complete(None, t2, result(10, 1000))
|
||||
t3 = Trial("t3", "PPO")
|
||||
t3 = Trial("PPO")
|
||||
self.assertEqual(
|
||||
rule.on_trial_result(None, t3, result(1, 260)),
|
||||
TrialScheduler.CONTINUE)
|
||||
@@ -120,8 +120,8 @@ class EarlyStoppingSuite(unittest.TestCase):
|
||||
rule = MedianStoppingRule(
|
||||
grace_period=0, min_samples_required=1,
|
||||
time_attr='training_iteration', reward_attr='neg_mean_loss')
|
||||
t1 = Trial("t1", "PPO") # mean is 450, max 900, t_max=10
|
||||
t2 = Trial("t2", "PPO") # mean is 450, max 450, t_max=5
|
||||
t1 = Trial("PPO") # mean is 450, max 900, t_max=10
|
||||
t2 = Trial("PPO") # mean is 450, max 450, t_max=5
|
||||
for i in range(10):
|
||||
self.assertEqual(
|
||||
rule.on_trial_result(None, t1, result2(i, i * 100)),
|
||||
@@ -166,7 +166,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
(81, 1) -> (27, 3) -> (9, 9) -> (3, 27) -> (1, 81);"""
|
||||
sched = HyperBandScheduler()
|
||||
for i in range(num_trials):
|
||||
t = Trial("t%d" % i, "__fake")
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
runner = _MockTrialRunner()
|
||||
return sched, runner
|
||||
@@ -211,7 +211,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
def advancedSetup(self):
|
||||
sched = self.basicSetup()
|
||||
for i in range(4):
|
||||
t = Trial("t%d" % (i + 20), "__fake")
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
|
||||
self.assertEqual(sched._cur_band_filled(), False)
|
||||
@@ -232,7 +232,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
sched = HyperBandScheduler()
|
||||
i = 0
|
||||
while not sched._cur_band_filled():
|
||||
t = Trial("t%d" % (i), "__fake")
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
i += 1
|
||||
self.assertEqual(len(sched._hyperbands[0]), 5)
|
||||
@@ -244,7 +244,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
sched = HyperBandScheduler(max_t=810)
|
||||
i = 0
|
||||
while not sched._cur_band_filled():
|
||||
t = Trial("t%d" % (i), "__fake")
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
i += 1
|
||||
self.assertEqual(len(sched._hyperbands[0]), 5)
|
||||
@@ -257,7 +257,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
sched = HyperBandScheduler(max_t=1)
|
||||
i = 0
|
||||
while len(sched._hyperbands) < 2:
|
||||
t = Trial("t%d" % (i), "__fake")
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
i += 1
|
||||
self.assertEqual(len(sched._hyperbands[0]), 5)
|
||||
@@ -415,7 +415,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
status = sched.on_trial_result(
|
||||
mock_runner, t, result(init_units, i))
|
||||
self.assertEqual(status, TrialScheduler.CONTINUE)
|
||||
t = Trial("t%d" % 100, "__fake")
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
mock_runner._launch_trial(t)
|
||||
self.assertEqual(len(sched._state["bracket"].current_trials()), 2)
|
||||
@@ -440,7 +440,7 @@ class HyperbandSuite(unittest.TestCase):
|
||||
stats = self.default_statistics()
|
||||
|
||||
for i in range(stats["max_trials"]):
|
||||
t = Trial("t%d" % i, "__fake")
|
||||
t = Trial("__fake")
|
||||
sched.on_trial_add(None, t)
|
||||
runner = _MockTrialRunner()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user