[rllib] Remove need to pass around registry (#2250)

* remove registry

* fix

* too many _

* fix

* cloudpickle

* Update registry.py

* yapf

* fix test

* fix kv check
This commit is contained in:
Eric Liang
2018-06-19 22:47:00 -07:00
committed by GitHub
parent 30684446a6
commit 30f7c08ca7
36 changed files with 202 additions and 208 deletions
+38 -41
View File
@@ -4,10 +4,10 @@ from __future__ import print_function
from types import FunctionType
import numpy as np
import ray
from ray.local_scheduler import ObjectID
import ray.cloudpickle as pickle
from ray.experimental.internal_kv import _internal_kv_initialized, \
_internal_kv_get, _internal_kv_put
TRAINABLE_CLASS = "trainable_class"
ENV_CREATOR = "env_creator"
@@ -35,7 +35,7 @@ def register_trainable(name, trainable):
if not issubclass(trainable, Trainable):
raise TypeError("Second argument must be convertable to Trainable",
trainable)
_default_registry.register(TRAINABLE_CLASS, name, trainable)
_global_registry.register(TRAINABLE_CLASS, name, trainable)
def register_env(name, env_creator):
@@ -48,62 +48,59 @@ def register_env(name, env_creator):
if not isinstance(env_creator, FunctionType):
raise TypeError("Second argument must be a function.", env_creator)
_default_registry.register(ENV_CREATOR, name, env_creator)
_global_registry.register(ENV_CREATOR, name, env_creator)
def get_registry():
"""Use this to access the registry. This requires ray to be initialized."""
def _make_key(category, key):
"""Generate a binary key for the given category and key.
_default_registry.flush_values_to_object_store()
Args:
category (str): The category of the item
key (str): The unique identifier for the item
# returns a registry copy that doesn't include the hard refs
return _Registry(_default_registry._all_objects)
def _to_pinnable(obj):
"""Converts obj to a form that can be pinned in object store memory.
Currently only numpy arrays are pinned in memory, if you have a strong
reference to the array value.
Returns:
The key to use for storing a the value.
"""
return (obj, np.zeros(1))
def _from_pinnable(obj):
"""Retrieve from _to_pinnable format."""
return obj[0]
return (b"TuneRegistry:" + category.encode("ascii") + b"/" +
key.encode("ascii"))
class _Registry(object):
def __init__(self, objs=None):
self._all_objects = {} if objs is None else objs.copy()
self._refs = [] # hard refs that prevent eviction of objects
def __init__(self):
self._to_flush = {}
def register(self, category, key, value):
if category not in KNOWN_CATEGORIES:
from ray.tune import TuneError
raise TuneError("Unknown category {} not among {}".format(
category, KNOWN_CATEGORIES))
self._all_objects[(category, key)] = value
self._to_flush[(category, key)] = pickle.dumps(value)
if _internal_kv_initialized():
self.flush_values()
def contains(self, category, key):
return (category, key) in self._all_objects
if _internal_kv_initialized():
value = _internal_kv_get(_make_key(category, key))
return value is not None
else:
return (category, key) in self._to_flush
def get(self, category, key):
value = self._all_objects[(category, key)]
if type(value) == ObjectID:
return _from_pinnable(ray.get(value))
if _internal_kv_initialized():
value = _internal_kv_get(_make_key(category, key))
if value is None:
raise ValueError(
"Registry value for {}/{} doesn't exist.".format(
category, key))
return pickle.loads(value)
else:
return value
return pickle.loads(self._to_flush[(category, key)])
def flush_values_to_object_store(self):
for k, v in self._all_objects.items():
if type(v) != ObjectID:
obj = ray.put(_to_pinnable(v))
self._all_objects[k] = obj
self._refs.append(ray.get(obj))
def flush_values(self):
for (category, key), value in self._to_flush.items():
_internal_kv_put(_make_key(category, key), value)
self._to_flush.clear()
_default_registry = _Registry()
_global_registry = _Registry()
ray.worker._post_init_hooks.append(_global_registry.flush_values)
+3 -3
View File
@@ -11,7 +11,7 @@ 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.registry import _global_registry, TRAINABLE_CLASS
from ray.tune.result import DEFAULT_RESULTS_DIR, TrainingResult
from ray.tune.util import pin_in_object_store, get_pinned_object
from ray.tune.experiment import Experiment
@@ -595,7 +595,7 @@ class TrialRunnerTest(unittest.TestCase):
def testTrialErrorOnStart(self):
ray.init()
_default_registry.register(TRAINABLE_CLASS, "asdf", None)
_global_registry.register(TRAINABLE_CLASS, "asdf", None)
trial = Trial("asdf", resources=Resources(1, 0))
try:
trial.start()
@@ -690,7 +690,7 @@ class TrialRunnerTest(unittest.TestCase):
},
"resources": Resources(cpu=1, gpu=1),
}
_default_registry.register(TRAINABLE_CLASS, "asdf", None)
_global_registry.register(TRAINABLE_CLASS, "asdf", None)
trials = [Trial("asdf", **kwargs), Trial("__fake", **kwargs)]
for t in trials:
runner.add_trial(t)
+1 -10
View File
@@ -46,11 +46,9 @@ class Trainable(object):
Attributes:
config (obj): The hyperparam configuration for this trial.
logdir (str): Directory in which training outputs should be placed.
registry (obj): Tune object registry which holds user-registered
classes and objects by name.
"""
def __init__(self, config=None, registry=None, logger_creator=None):
def __init__(self, config=None, logger_creator=None):
"""Initialize an Trainable.
Subclasses should prefer defining ``_setup()`` instead of overriding
@@ -58,20 +56,13 @@ class Trainable(object):
Args:
config (dict): Trainable-specific configuration data.
registry (obj): Object registry for user-defined envs, models, etc.
If unspecified, the default registry will be used.
logger_creator (func): Function that creates a ray.tune.Logger
object. If unspecified, a default logger is created.
"""
if registry is None:
from ray.tune.registry import get_registry
registry = get_registry()
self._initialize_ok = False
self._experiment_id = uuid.uuid4().hex
self.config = config or {}
self.registry = registry
if logger_creator:
self._result_logger = logger_creator(self.config)
+3 -5
View File
@@ -57,7 +57,7 @@ class Resources(
def has_trainable(trainable_name):
return ray.tune.registry._default_registry.contains(
return ray.tune.registry._global_registry.contains(
ray.tune.registry.TRAINABLE_CLASS, trainable_name)
@@ -377,12 +377,10 @@ class Trial(object):
# Logging for trials is handled centrally by TrialRunner, so
# configure the remote runner to use a noop-logger.
self.runner = cls.remote(
config=self.config,
registry=ray.tune.registry.get_registry(),
logger_creator=logger_creator)
config=self.config, logger_creator=logger_creator)
def _get_trainable_cls(self):
return ray.tune.registry.get_registry().get(
return ray.tune.registry._global_registry.get(
ray.tune.registry.TRAINABLE_CLASS, self.trainable_name)
def set_verbose(self, verbose):
+18 -2
View File
@@ -2,12 +2,12 @@ from __future__ import absolute_import
from __future__ import division
from __future__ import print_function
import base64
from six.moves import queue
import base64
import numpy as np
import threading
import ray
from ray.tune.registry import _to_pinnable, _from_pinnable
_pinned_objects = []
_fetch_requests = queue.Queue()
@@ -63,6 +63,22 @@ def _serve_get_pin_requests():
pass
def _to_pinnable(obj):
"""Converts obj to a form that can be pinned in object store memory.
Currently only numpy arrays are pinned in memory, if you have a strong
reference to the array value.
"""
return (obj, np.zeros(1))
def _from_pinnable(obj):
"""Retrieve from _to_pinnable format."""
return obj[0]
if __name__ == '__main__':
ray.init()
X = pin_in_object_store("hello")