mirror of
https://github.com/wassname/ray.git
synced 2026-07-25 13:30:52 +08:00
[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:
+38
-41
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user