mirror of
https://github.com/wassname/ray.git
synced 2026-07-23 13:10:11 +08:00
[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:
@@ -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),
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user