From 5f4a9cc713eec81d251c9abd1080957913ab3b29 Mon Sep 17 00:00:00 2001 From: Eric Liang Date: Tue, 11 Dec 2018 20:16:38 -0800 Subject: [PATCH] [rllib] Rollout should preprocess observations; some cleanups (#3512) ## What do these changes do? From https://groups.google.com/forum/#!topic/ray-dev/u-gybKK6-Ns --- python/ray/rllib/agents/agent.py | 4 +- python/ray/rllib/models/catalog.py | 63 --------------------- python/ray/rllib/rollout.py | 24 +++++--- python/ray/rllib/test/test_nested_spaces.py | 16 ++++++ 4 files changed, 35 insertions(+), 72 deletions(-) diff --git a/python/ray/rllib/agents/agent.py b/python/ray/rllib/agents/agent.py index 8e6797eed..5f0631e11 100644 --- a/python/ray/rllib/agents/agent.py +++ b/python/ray/rllib/agents/agent.py @@ -315,8 +315,10 @@ class Agent(Trainable): if state is None: state = [] + preprocessed = self.local_evaluator.preprocessors[policy_id].transform( + observation) filtered_obs = self.local_evaluator.filters[policy_id]( - observation, update=False) + preprocessed, update=False) if state: return self.local_evaluator.for_policy( lambda p: p.compute_single_action(filtered_obs, state), diff --git a/python/ray/rllib/models/catalog.py b/python/ray/rllib/models/catalog.py index 822af4a37..b3a84cef3 100644 --- a/python/ray/rllib/models/catalog.py +++ b/python/ray/rllib/models/catalog.py @@ -11,9 +11,6 @@ from functools import partial from ray.tune.registry import RLLIB_MODEL, RLLIB_PREPROCESSOR, \ _global_registry -from ray.rllib.env.async_vector_env import _ExternalEnvToAsync -from ray.rllib.env.external_env import ExternalEnv -from ray.rllib.env.vector_env import VectorEnv from ray.rllib.models.action_dist import ( Categorical, Deterministic, DiagGaussian, MultiActionDistribution) from ray.rllib.models.preprocessors import get_preprocessor @@ -310,29 +307,6 @@ class ModelCatalog(object): prep, observation_space, prep.shape)) return prep - @staticmethod - def get_preprocessor_as_wrapper(env, options=None): - """Returns a preprocessor as a gym observation wrapper. - - Args: - env (gym.Env|VectorEnv|ExternalEnv): The environment to wrap. - options (dict): Options to pass to the preprocessor. - - Returns: - env (RLlib env): Wrapped environment - """ - - options = options or MODEL_DEFAULTS - preprocessor = ModelCatalog.get_preprocessor(env, options) - if isinstance(env, gym.Env): - return _RLlibPreprocessorWrapper(env, preprocessor) - elif isinstance(env, VectorEnv): - return _RLlibVectorPreprocessorWrapper(env, preprocessor) - elif isinstance(env, ExternalEnv): - return _ExternalEnvToAsync(env, preprocessor) - else: - raise ValueError("Don't know how to wrap {}".format(env)) - @staticmethod def register_custom_preprocessor(preprocessor_name, preprocessor_class): """Register a custom preprocessor class by name. @@ -359,40 +333,3 @@ class ModelCatalog(object): model_class (type): Python class of the model. """ _global_registry.register(RLLIB_MODEL, model_name, model_class) - - -class _RLlibPreprocessorWrapper(gym.ObservationWrapper): - """Adapts a RLlib preprocessor for use as an observation wrapper.""" - - def __init__(self, env, preprocessor): - super(_RLlibPreprocessorWrapper, self).__init__(env) - self.preprocessor = preprocessor - self.observation_space = preprocessor.observation_space - - def observation(self, observation): - return self.preprocessor.transform(observation) - - -class _RLlibVectorPreprocessorWrapper(VectorEnv): - """Preprocessing wrapper for vector envs.""" - - def __init__(self, env, preprocessor): - self.env = env - self.prep = preprocessor - self.action_space = env.action_space - self.observation_space = preprocessor.observation_space - self.num_envs = env.num_envs - - def vector_reset(self): - return [self.prep.transform(obs) for obs in self.env.vector_reset()] - - def reset_at(self, index): - return self.prep.transform(self.env.reset_at(index)) - - def vector_step(self, actions): - obs, rewards, dones, infos = self.env.vector_step(actions) - obs = [self.prep.transform(o) for o in obs] - return obs, rewards, dones, infos - - def get_unwrapped(self): - return self.env.get_unwrapped() diff --git a/python/ray/rllib/rollout.py b/python/ray/rllib/rollout.py index bee5c5eb2..01cd0b915 100755 --- a/python/ray/rllib/rollout.py +++ b/python/ray/rllib/rollout.py @@ -23,6 +23,11 @@ Example Usage via executable: --env CartPole-v0 --steps 1000000 --out rollouts.pkl """ +# Note: if you use any custom models or envs, register them here first, e.g.: +# +# ModelCatalog.register_custom_model("pa_model", ParametricActionsModel) +# register_env("pa_cartpole", lambda _: ParametricActionCartpole(10)) + def create_parser(parser_creator=None): parser_creator = parser_creator or argparse.ArgumentParser @@ -91,16 +96,19 @@ def run(args, parser): agent = cls(env=args.env, config=config) agent.restore(args.checkpoint) num_steps = int(args.steps) + rollout(agent, args.env, num_steps, args.out, args.no_render) + +def rollout(agent, env_name, num_steps, out=None, no_render=True): if hasattr(agent, "local_evaluator"): env = agent.local_evaluator.env else: - env = gym.make(args.env) - if args.out is not None: + env = gym.make(env_name) + if out is not None: rollouts = [] steps = 0 while steps < (num_steps or steps + 1): - if args.out is not None: + if out is not None: rollout = [] state = env.reset() done = False @@ -109,17 +117,17 @@ def run(args, parser): action = agent.compute_action(state) next_state, reward, done, _ = env.step(action) reward_total += reward - if not args.no_render: + if not no_render: env.render() - if args.out is not None: + if out is not None: rollout.append([state, action, next_state, reward, done]) steps += 1 state = next_state - if args.out is not None: + if out is not None: rollouts.append(rollout) print("Episode reward", reward_total) - if args.out is not None: - pickle.dump(rollouts, open(args.out, "wb")) + if out is not None: + pickle.dump(rollouts, open(out, "wb")) if __name__ == "__main__": diff --git a/python/ray/rllib/test/test_nested_spaces.py b/python/ray/rllib/test/test_nested_spaces.py index 95744b7e2..90e527aa4 100644 --- a/python/ray/rllib/test/test_nested_spaces.py +++ b/python/ray/rllib/test/test_nested_spaces.py @@ -19,6 +19,7 @@ from ray.rllib.env.async_vector_env import AsyncVectorEnv from ray.rllib.env.vector_env import VectorEnv from ray.rllib.models import ModelCatalog from ray.rllib.models.model import Model +from ray.rllib.rollout import rollout from ray.rllib.test.test_external_env import SimpleServing from ray.tune.registry import register_env @@ -340,6 +341,21 @@ class NestedSpacesTest(unittest.TestCase): self.assertEqual(seen[1][0].tolist(), cam_i) self.assertEqual(seen[2][0].tolist(), task_i) + def testRolloutDictSpace(self): + register_env("nested", lambda _: NestedDictEnv()) + agent = PGAgent(env="nested") + agent.train() + path = agent.save() + agent.stop() + + # Test train works on restore + agent2 = PGAgent(env="nested") + agent2.restore(path) + agent2.train() + + # Test rollout works on restore + rollout(agent2, "nested", 100) + if __name__ == "__main__": ray.init(num_cpus=5)