[rllib] Rollout should preprocess observations; some cleanups (#3512)

<!--
Thank you for your contribution!

Please review https://github.com/ray-project/ray/blob/master/CONTRIBUTING.rst before opening a pull request.
-->

## What do these changes do?

From https://groups.google.com/forum/#!topic/ray-dev/u-gybKK6-Ns
This commit is contained in:
Eric Liang
2018-12-11 20:16:38 -08:00
committed by Richard Liaw
parent 59f4743f20
commit 5f4a9cc713
4 changed files with 35 additions and 72 deletions
+3 -1
View File
@@ -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),
-63
View File
@@ -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()
+16 -8
View File
@@ -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__":
@@ -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)