[tune] Initial Commit for Tune CLI (#3983)

This introduces a light CLI for Tune.
This commit is contained in:
Richard Liaw
2019-03-08 16:46:05 -08:00
committed by GitHub
parent 3064fad96b
commit 6630a35353
18 changed files with 492 additions and 69 deletions
+15 -14
View File
@@ -25,7 +25,8 @@ from ray.tune.logger import pretty_print, UnifiedLogger
# have been defined yet. See https://github.com/ray-project/ray/issues/1716.
import ray.tune.registry
from ray.tune.result import (DEFAULT_RESULTS_DIR, DONE, HOSTNAME, PID,
TIME_TOTAL_S, TRAINING_ITERATION, TIMESTEPS_TOTAL)
TIME_TOTAL_S, TRAINING_ITERATION, TIMESTEPS_TOTAL,
EPISODE_REWARD_MEAN, MEAN_LOSS, MEAN_ACCURACY)
from ray.utils import _random_string, binary_to_hex, hex_to_binary
DEBUG_PRINT_INTERVAL = 5
@@ -299,6 +300,8 @@ class Trial(object):
self.error_file = None
self.num_failures = 0
self.custom_trial_name = None
# AutoML fields
self.results = None
self.best_result = None
@@ -316,10 +319,8 @@ class Trial(object):
"param_config",
"extra_arg",
]
self.trial_name = None
if trial_name_creator:
self.trial_name = trial_name_creator(self)
self.custom_trial_name = trial_name_creator(self)
@classmethod
def _registration_check(cls, trainable_name):
@@ -447,17 +448,17 @@ class Trial(object):
if self.last_result.get(TIMESTEPS_TOTAL) is not None:
pieces.append('{} ts'.format(self.last_result[TIMESTEPS_TOTAL]))
if self.last_result.get("episode_reward_mean") is not None:
if self.last_result.get(EPISODE_REWARD_MEAN) is not None:
pieces.append('{} rew'.format(
format(self.last_result["episode_reward_mean"], '.3g')))
format(self.last_result[EPISODE_REWARD_MEAN], '.3g')))
if self.last_result.get("mean_loss") is not None:
if self.last_result.get(MEAN_LOSS) is not None:
pieces.append('{} loss'.format(
format(self.last_result["mean_loss"], '.3g')))
format(self.last_result[MEAN_LOSS], '.3g')))
if self.last_result.get("mean_accuracy") is not None:
if self.last_result.get(MEAN_ACCURACY) is not None:
pieces.append('{} acc'.format(
format(self.last_result["mean_accuracy"], '.3g')))
format(self.last_result[MEAN_ACCURACY], '.3g')))
return ', '.join(pieces)
@@ -514,8 +515,8 @@ class Trial(object):
Can be overriden with a custom string creator.
"""
if self.trial_name:
return self.trial_name
if self.custom_trial_name:
return self.custom_trial_name
if "env" in self.config:
env = self.config["env"]
@@ -544,8 +545,6 @@ class Trial(object):
state["runner"] = None
state["result_logger"] = None
if self.status == Trial.RUNNING:
state["status"] = Trial.PENDING
if self.result_logger:
self.result_logger.flush()
state["__logger_started__"] = True
@@ -556,6 +555,8 @@ class Trial(object):
def __setstate__(self, state):
logger_started = state.pop("__logger_started__")
state["resources"] = json_to_resources(state["resources"])
if state["status"] == Trial.RUNNING:
state["status"] = Trial.PENDING
for key in self._nonjson_fields:
state[key] = cloudpickle.loads(hex_to_binary(state[key]))