A3C Polishing (#385)

* number

* gym doesn't have versioning

* Benchmarks

* visualization

* formatting

* small fix for tensorboard

* first pass removing universe dependency

* code

* results polish

* removed extra line

* removed universe dependency

* doc

* remove gym versioning stuff

* changes as suggested

* nit
This commit is contained in:
Richard Liaw
2017-04-11 22:51:52 -07:00
committed by Robert Nishihara
parent c3a2505ffd
commit 94f32db5e6
11 changed files with 59 additions and 700 deletions
+9 -3
View File
@@ -17,12 +17,10 @@ To run the application, first install **ray** and then some dependencies:
pip install tensorflow
pip install six
pip install gym[atari]==0.7.4
pip install gym[atari]
pip install opencv-python
pip install scipy
Note that this code **currently does not work** with ``gym==0.8.0``.
You can run the code with
.. code-block:: bash
@@ -144,3 +142,11 @@ global model parameters. The main training script looks like the following.
parameters = policy.get_weights()
gradient_list.extend([agents[info["id"]].compute_gradient(parameters)])
return policy
Benchmarks and Visualization
----------------------------
For the :code:`PongDeterministic-v3` and an Amazon EC2 m4.16xlarge instance, we are able to train the agent with 16 workers in around 15 minutes. With 8 workers, we can train the agent in around 25 minutes.
You can visualize performance by running :code:`tensorboard --logdir [directory]` in a separate screen, where :code:`[directory]` is defaulted to :code:`./results/`. If you are running multiple experiments, be sure to vary the directory to which Tensorflow saves its progress (found in :code:`driver.py`).
+6 -8
View File
@@ -10,16 +10,17 @@ import tensorflow as tf
import six.moves.queue as queue
import gym
import sys
import os
from datetime import datetime, timedelta
from misc import timestamp, Profiler
from misc import timestamp, time_string
from envs import create_env
@ray.actor
class Runner(object):
"""Actor object to start running simulation on workers.
Gradient computation is also executed from this object."""
def __init__(self, env_name, actor_id, logdir="tmp/", start=True):
env = create_env(env_name, None, None)
def __init__(self, env_name, actor_id, logdir="results/", start=True):
env = create_env(env_name)
self.id = actor_id
num_actions = env.action_space.n
self.policy = LSTMPolicy(env.observation_space.shape, num_actions, actor_id)
@@ -40,7 +41,7 @@ class Runner(object):
return rollout
def start(self):
summary_writer = tf.summary.FileWriter(self.logdir + "test_1")
summary_writer = tf.summary.FileWriter(os.path.join(self.logdir, "agent_%d" % self.id))
self.summary_writer = summary_writer
self.runner.start_runner(self.policy.sess, summary_writer)
@@ -55,7 +56,7 @@ class Runner(object):
def train(num_workers, env_name="PongDeterministic-v3"):
env = create_env(env_name, None, None)
env = create_env(env_name)
policy = LSTMPolicy(env.observation_space.shape, env.action_space.n, 0)
agents = [Runner(env_name, i) for i in range(num_workers)]
parameters = policy.get_weights()
@@ -73,9 +74,6 @@ def train(num_workers, env_name="PongDeterministic-v3"):
return policy
if __name__ == '__main__':
if gym.__version__[:3] == '0.8':
raise Exception("This example currently does not work with gym==0.8.0. "
"Please downgrade to gym==0.7.4.");
num_workers = int(sys.argv[1])
ray.init(num_cpus=num_workers)
train(num_workers)
+38 -97
View File
@@ -10,29 +10,52 @@ import logging
import numpy as np
import time
import vectorized
from vectorized.wrappers import Unvectorize, Vectorize
logger = logging.getLogger(__name__)
logger.setLevel(logging.INFO)
def create_env(env_id, client_id, remotes, **kwargs):
return create_atari_env(env_id)
def create_atari_env(env_id):
def create_env(env_id):
env = gym.make(env_id)
env = Vectorize(env)
env = AtariRescale42x42(env)
env = DiagnosticsInfo(env)
env = Unvectorize(env)
env = AtariProcessing(env)
env = Diagnostic(env)
return env
def DiagnosticsInfo(env, *args, **kwargs):
return vectorized.VectorizeFilter(env, DiagnosticsInfoI, *args, **kwargs)
def _process_frame42(frame):
frame = frame[34:(34+160), :160]
# Resize by half, then down to 42x42 (essentially mipmapping). If
# we resize directly we lose pixels that, when mapped to 42x42,
# aren't close enough to the pixel boundary.
frame = cv2.resize(frame, (80, 80))
frame = cv2.resize(frame, (42, 42))
frame = frame.mean(2)
frame = frame.astype(np.float32)
frame *= (1.0 / 255.0)
frame = np.reshape(frame, [42, 42, 1])
return frame
class DiagnosticsInfoI(vectorized.Filter):
class AtariProcessing(gym.ObservationWrapper):
def __init__(self, env=None):
super(AtariProcessing, self).__init__(env)
self.observation_space = Box(0.0, 1.0, [42, 42, 1])
def _observation(self, observation):
return _process_frame42(observation)
class Diagnostic(gym.Wrapper):
def __init__(self, env=None):
super(Diagnostic, self).__init__(env)
self.diagnostics = DiagnosticsLogger()
def _reset(self):
observation = self.env.reset()
return self.diagnostics._after_reset(observation)
def _step(self, action):
results = self.env.step(action)
return self.diagnostics._after_step(*results)
class DiagnosticsLogger():
def __init__(self, log_interval=503):
super(DiagnosticsInfoI, self).__init__()
self._episode_time = time.time()
self._last_time = time.time()
@@ -41,7 +64,6 @@ class DiagnosticsInfoI(vectorized.Filter):
self._episode_reward = 0
self._episode_length = 0
self._all_rewards = []
self._num_vnc_updates = 0
self._last_episode_id = -1
def _after_reset(self, observation):
@@ -57,43 +79,12 @@ class DiagnosticsInfoI(vectorized.Filter):
self._episode_time = time.time()
self._local_t += 1
if info.get("stats.vnc.updates.n") is not None:
self._num_vnc_updates += info.get("stats.vnc.updates.n")
if self._local_t % self._log_interval == 0:
cur_time = time.time()
elapsed = cur_time - self._last_time
fps = self._log_interval / elapsed
self._last_time = cur_time
cur_episode_id = info.get('vectorized.episode_id', 0)
to_log["diagnostics/fps"] = fps
if self._last_episode_id == cur_episode_id:
to_log["diagnostics/fps_within_episode"] = fps
self._last_episode_id = cur_episode_id
if info.get("stats.gauges.diagnostics.lag.action") is not None:
to_log["diagnostics/action_lag_lb"] = info["stats.gauges.diagnostics.lag.action"][0]
to_log["diagnostics/action_lag_ub"] = info["stats.gauges.diagnostics.lag.action"][1]
if info.get("reward.count") is not None:
to_log["diagnostics/reward_count"] = info["reward.count"]
if info.get("stats.gauges.diagnostics.clock_skew") is not None:
to_log["diagnostics/clock_skew_lb"] = info["stats.gauges.diagnostics.clock_skew"][0]
to_log["diagnostics/clock_skew_ub"] = info["stats.gauges.diagnostics.clock_skew"][1]
if info.get("stats.gauges.diagnostics.lag.observation") is not None:
to_log["diagnostics/observation_lag_lb"] = info["stats.gauges.diagnostics.lag.observation"][0]
to_log["diagnostics/observation_lag_ub"] = info["stats.gauges.diagnostics.lag.observation"][1]
if info.get("stats.vnc.updates.n") is not None:
to_log["diagnostics/vnc_updates_n"] = info["stats.vnc.updates.n"]
to_log["diagnostics/vnc_updates_n_ps"] = self._num_vnc_updates / elapsed
self._num_vnc_updates = 0
if info.get("stats.vnc.updates.bytes") is not None:
to_log["diagnostics/vnc_updates_bytes"] = info["stats.vnc.updates.bytes"]
if info.get("stats.vnc.updates.pixels") is not None:
to_log["diagnostics/vnc_updates_pixels"] = info["stats.vnc.updates.pixels"]
if info.get("stats.vnc.updates.rectangles") is not None:
to_log["diagnostics/vnc_updates_rectangles"] = info["stats.vnc.updates.rectangles"]
if info.get("env_status.state_id") is not None:
to_log["diagnostics/env_state_id"] = info["env_status.state_id"]
if reward is not None:
self._episode_reward += reward
@@ -114,53 +105,3 @@ class DiagnosticsInfoI(vectorized.Filter):
return observation, reward, done, to_log
def _process_frame42(frame):
frame = frame[34:34+160, :160]
# Resize by half, then down to 42x42 (essentially mipmapping). If
# we resize directly we lose pixels that, when mapped to 42x42,
# aren't close enough to the pixel boundary.
frame = cv2.resize(frame, (80, 80))
frame = cv2.resize(frame, (42, 42))
frame = frame.mean(2)
frame = frame.astype(np.float32)
frame *= (1.0 / 255.0)
frame = np.reshape(frame, [42, 42, 1])
return frame
class AtariRescale42x42(vectorized.ObservationWrapper):
def __init__(self, env=None):
super(AtariRescale42x42, self).__init__(env)
self.observation_space = Box(0.0, 1.0, [42, 42, 1])
def _observation(self, observation_n):
return [_process_frame42(observation) for observation in observation_n]
class CropScreen(vectorized.ObservationWrapper):
"""Crops out a [height]x[width] area starting from (top,left) """
def __init__(self, env, height, width, top=0, left=0):
super(CropScreen, self).__init__(env)
self.height = height
self.width = width
self.top = top
self.left = left
self.observation_space = Box(0, 255, shape=(height, width, 3))
def _observation(self, observation_n):
return [ob[self.top:self.top+self.height, self.left:self.left+self.width, :] if ob is not None else None
for ob in observation_n]
def _process_frame_flash(frame):
frame = cv2.resize(frame, (200, 128))
frame = frame.mean(2).astype(np.float32)
frame *= (1.0 / 255.0)
frame = np.reshape(frame, [128, 200, 1])
return frame
class FlashRescale(vectorized.ObservationWrapper):
def __init__(self, env=None):
super(FlashRescale, self).__init__(env)
self.observation_space = Box(0.0, 1.0, [128, 200, 1])
def _observation(self, observation_n):
return [_process_frame_flash(observation) for observation in observation_n]
+3
View File
@@ -9,6 +9,9 @@ import cProfile, pstats, io
def timestamp():
return datetime.now().timestamp()
def time_string():
return datetime.now().strftime("%Y%m%d_%H_%M_%f")
class Profiler(object):
def __init__(self):
self.pr = cProfile.Profile()
+3 -2
View File
@@ -113,6 +113,7 @@ runner appends the policy to the queue.
last_features = policy.get_initial_features()
length = 0
rewards = 0
rollout_number = 0
while True:
terminal_end = False
@@ -138,7 +139,7 @@ runner appends the policy to the queue.
summary = tf.Summary()
for k, v in info.items():
summary.value.add(tag=k, simple_value=float(v))
summary_writer.add_summary(summary, policy.global_step.eval())
summary_writer.add_summary(summary, rollout_number)
summary_writer.flush()
timestep_limit = env.spec.tags.get('wrapper_config.TimeLimit.max_episode_steps')
@@ -147,7 +148,7 @@ runner appends the policy to the queue.
if length >= timestep_limit or not env.metadata.get('semantics.autoreset'):
last_state = env.reset()
last_features = policy.get_initial_features()
# print("Episode finished. Sum of rewards: %d. Length: %d" % (rewards, length))
rollout_number += 1
length = 0
rewards = 0
break
-8
View File
@@ -1,8 +0,0 @@
from __future__ import absolute_import
from __future__ import division
from __future__ import print_function
from .vectorize_core import Env, Wrapper, ObservationWrapper, ActionWrapper, RewardWrapper
from .multiprocessing_env import MultiprocessingEnv
from .vectorize_filter import Filter, VectorizeFilter
from .wrappers import Vectorize, Unvectorize, WeakUnvectorize
-56
View File
@@ -1,56 +0,0 @@
from __future__ import absolute_import
from __future__ import division
from __future__ import print_function
import weakref
from gym import monitoring
class Monitor(object):
def __init__(self, env_n):
"""env_n is a collection of unvectorized envs"""
self.monitor_n = [monitoring.Monitor(env) for env in env_n]
@property
def env(self):
# The real env is the first unwrapped env. Maybe we should
# maintain our own weakref rather than doing this.
return self.monitor_n[0].env.env
def start(self, directory, video_callable=None, seed_n=None, force=False,
resume=False, write_upon_reset=False, uid=None):
if seed_n is None:
seed_n = [None] * len(self.monitor_n)
# There's way to seed just one of the vectorized environments,
# so we have to do the seeding ourselves outside of the
# underlying monitor instances.
#
# The monitor will call the .seed method on the
# WeakUnvectorized env, which just returns rather than
# actually re-seeding the env.
self.env.seed(seed_n)
for i, monitor in enumerate(self.monitor_n):
# Only allow recording of video in first monitor
if i > 0:
video_callable = False
# Seed gets passed in but just recorded, not used.
monitor.start(directory=directory, video_callable=video_callable,
force=force, resume=resume, write_upon_reset=write_upon_reset, uid=uid)
def close(self, *args, **kwargs):
[monitor.close(*args, **kwargs) for monitor in self.monitor_n]
def _before_reset(self):
return [monitor._before_reset() for monitor in self.monitor_n]
def _after_reset(self, observation_n):
assert len(observation_n) == len(self.monitor_n)
return [monitor._after_reset(observation) for monitor, observation in zip(self.monitor_n, observation_n)]
def _before_step(self, action_n):
assert len(action_n) == len(self.monitor_n)
return [monitor._before_step(action) for monitor, action in zip(self.monitor_n, action_n)]
def _after_step(self, observation_n, reward_n, done_n, info):
return [monitor._after_step(o, r, d, i) for monitor, o, r, d, i in zip(self.monitor_n, observation_n, reward_n, done_n, info['n'])]
@@ -1,326 +0,0 @@
from __future__ import absolute_import
from __future__ import division
from __future__ import print_function
import logging
import multiprocessing
import numpy as np
import traceback
import gym
from gym import spaces
import vectorized.vectorize_core as core
logger = logging.getLogger(__name__)
logger.setLevel(logging.INFO)
class Error(Exception):
pass
def display_name(exception):
prefix = ''
# AttributeError has no __module__; RuntimeError has module of
# exceptions
if hasattr(exception, '__module__') and exception.__module__ != 'exceptions':
prefix = exception.__module__ + '.'
return prefix + type(exception).__name__
def render_dict(error):
return {
'type': display_name(error),
'message': error.message,
'traceback': traceback.format_exc(error)
}
class Worker(object):
def __init__(self, env_m, worker_idx):
# These are instantiated in the *parent* process
# currently. Probably will want to change this. The parent
# does need to obtain the relevant Spaces at some stage, but
# that's doable.
self.worker_idx = worker_idx
self.env_m = env_m
self.m = len(env_m)
self.parent_conn, self.child_conn = multiprocessing.Pipe()
self.joiner = multiprocessing.Process(target=self.run)
self._clear_state()
self.start()
# Parent only!
self.child_conn.close()
def _clear_state(self):
self.mask = [True] * self.m
# Control methods
def start(self):
self.joiner.start()
def _parent_recv(self):
rendered, res = self.parent_conn.recv()
if rendered is not None:
raise Error('[Worker {}] Error: {} ({})\n\n{}'.format(self.worker_idx, rendered['message'], rendered['type'], rendered['traceback']))
return res
def _child_send(self, msg):
self.child_conn.send((None, msg))
def _parent_send(self, msg):
try:
self.parent_conn.send(msg)
except IOError: # the worker is now dead
try:
res = self._parent_recv()
except EOFError:
raise Error('[Worker {}] Child died unexpectedly'.format(self.worker_idx))
else:
raise Error('[Worker {}] Child returned unexpected result: {}'.format(self.worker_idx, res))
def close_start(self):
self._parent_send(('close', None))
def close_finish(self):
self.joiner.join()
def reset_start(self):
self._parent_send(('reset', None))
def reset_finish(self):
return self._parent_recv()
def step_start(self, action_m):
"""action_m: the batch of actions for this worker"""
self._parent_send(('step', action_m))
def step_finish(self):
return self._parent_recv()
def mask_start(self, i):
self._parent_send(('mask', i))
def seed_start(self, seed_m):
self._parent_send(('seed', seed_m))
def render_start(self, mode, close):
self._parent_send(('render', (mode, close)))
def render_finish(self):
return self._parent_recv()
def run(self):
try:
self.do_run()
except Exception as e:
rendered = render_dict(e)
self.child_conn.send((rendered, None))
return
def do_run(self):
# Child only!
self.parent_conn.close()
while True:
method, body = self.child_conn.recv()
logger.debug('[%d] Received: method=%s body=%s', self.worker_idx, method, body)
if method == 'close':
logger.info('Closing envs')
# TODO: close envs?
return
elif method == 'reset':
self._clear_state()
observation_m = [env.reset() for env in self.env_m]
self._child_send(observation_m)
elif method == 'step':
action_m = body
observation_m, reward_m, done_m, info = self.step_m(action_m)
self._child_send((observation_m, reward_m, done_m, info))
elif method == 'mask':
i = body
assert 0 <= i < self.m, 'Bad value for mask: {} (should be >= 0 and < {})'.format(i, self.m)
self.mask[i] = False
logger.debug('[%d] Applying mask: i=%d', self.worker_idx, i)
elif method == 'seed':
seeds = body
[env.seed(seed) for env, seed in zip(self.env_m, seeds)]
elif method == 'render':
mode, close = body
if mode == 'human':
self.env_m[0].render(mode=mode, close=close)
result = [None]
else:
result = [env.render(mode=mode, close=close) for env in self.env_m]
self._child_send(result)
else:
raise Error('Bad method: {}'.format(method))
def step_m(self, action_m):
observation_m = []
reward_m = []
done_m = []
info = {'m': []}
for env, enabled, action in zip(self.env_m, self.mask, action_m):
if enabled:
observation, reward, done, info_i = env.step(action)
if done:
observation = env.reset()
else:
observation = None
reward = 0
done = False
info_i = {}
observation_m.append(observation)
reward_m.append(reward)
done_m.append(done)
info['m'].append(info_i)
return observation_m, reward_m, done_m, info
def step_n(worker_n, action_n):
accumulated = 0
for worker in worker_n:
action_m = action_n[accumulated:accumulated+worker.m]
worker.step_start(action_m)
accumulated += worker.m
observation_n = []
reward_n = []
done_n = []
info = {'n': []}
for worker in worker_n:
observation_m, reward_m, done_m, info_i = worker.step_finish()
observation_n += observation_m
reward_n += reward_m
done_n += done_m
info['n'] += info_i['m']
return observation_n, reward_n, done_n, info
def reset_n(worker_n):
for worker in worker_n:
worker.reset_start()
observation_n = []
for worker in worker_n:
observation_n += worker.reset_finish()
return observation_n
def seed_n(worker_n, seed_n):
accumulated = 0
for worker in worker_n:
action_m = seed_n[accumulated:accumulated+worker.m]
worker.seed_start(seed_n)
accumulated += worker.m
def mask(worker_n, i):
accumulated = 0
for k, worker in enumerate(worker_n):
if accumulated + worker.m <= i:
accumulated += worker.m
else:
worker.mask_start(i - accumulated)
return
def render_n(worker_n, mode, close):
if mode == 'human':
# Only render 1 worker
worker_n = worker_n[0:]
for worker in worker_n:
worker.render_start(mode, close)
res = []
for worker in worker_n:
res += worker.render_finish()
if mode != 'human':
return res
else:
return None
def close_n(worker_n):
if worker_n is None:
return
# TODO: better error handling: workers should die when we go away
# anyway. Also technically should wait for these processes if
# we're not crashing.
for worker in worker_n:
try:
worker.close_start()
except Error:
pass
# for worker in worker_n:
# try:
# worker.close_finish()
# except Error:
# pass
class MultiprocessingEnv(core.Env):
metadata = {
'runtime.vectorized': True,
}
def __init__(self, env_id):
self.worker_n = None
# Pull the relevant info from a transient env instance
self.spec = gym.spec(env_id)
env = self.spec.make()
current_metadata = self.metadata
self.metadata = env.metadata.copy()
self.metadata.update(current_metadata)
self.action_space = env.action_space
self.observation_space = env.observation_space
self.reward_range = env.reward_range
def _configure(self, n=1, pool_size=None, episode_limit=None):
super(MultiprocessingEnv, self)._configure()
self.n = n
self.envs = [self.spec.make() for _ in range(self.n)]
if pool_size is None:
pool_size = min(len(self.envs), multiprocessing.cpu_count() - 1)
pool_size = max(1, pool_size)
self.worker_n = []
m = int((self.n + pool_size - 1) / pool_size)
for i in range(0, self.n, m):
envs = self.envs[i:i+m]
self.worker_n.append(Worker(envs, i))
if episode_limit is not None:
self._episode_id.episode_limit = episode_limit
def _seed(self, seed):
seed_n(self.worker_n, seed)
return [[seed_i] for seed_i in seed]
def _reset(self):
return reset_n(self.worker_n)
def _step(self, action_n):
return step_n(self.worker_n, action_n)
def _render(self, mode='human', close=False):
return render_n(self.worker_n, mode=mode, close=close)
def mask(self, i):
mask(self.worker_n, i)
def _close(self):
close_n(self.worker_n)
if __name__ == '__main__':
env_n = make('Pong-v3')
env_n.configure()
env_n.reset()
print(env_n.step([0] * 10))
-53
View File
@@ -1,53 +0,0 @@
from __future__ import absolute_import
from __future__ import division
from __future__ import print_function
import gym
from gym import spaces
class Env(gym.Env):
"""Base class capable of handling vectorized environments.
"""
metadata = {
# This key indicates whether an env is vectorized (or, in the case of
# Wrappers where autovectorize=True, whether they should automatically
# be wrapped by a Vectorize wrapper.)
'runtime.vectorized': True,
}
# Number of remotes. User should set this.
n = None
class Wrapper(Env, gym.Wrapper):
"""Use this instead of gym.Wrapper iff you're wrapping a vectorized env,
(or a vanilla env you wish to be vectorized).
"""
# If True and this is instantiated with a non-vectorized environment,
# automatically wrap it with the Vectorize wrapper.
autovectorize = True
def __init__(self, env):
super(Wrapper, self).__init__(env)
if not env.metadata.get('runtime.vectorized'):
if self.autovectorize:
# Circular dependency :(
import vectorize.wrappers as wrappers
env = wrappers.Vectorize(env)
else:
raise Exception('This wrapper can only wrap vectorized envs (i.e. where env.metadata["runtime.vectorized"] = True), not {}. Set "self.autovectorize = True" to automatically add a Vectorize wrapper.'.format(env))
self.env = env
def _configure(self, **kwargs):
super(Wrapper, self)._configure(**kwargs)
assert self.env.n is not None, "Did not set self.env.n: self.n={} self.env={} self={}".format(self.env.n, self.env, self)
self.n = self.env.n
class ObservationWrapper(Wrapper, gym.ObservationWrapper):
pass
class RewardWrapper(Wrapper, gym.RewardWrapper):
pass
class ActionWrapper(Wrapper, gym.ActionWrapper):
pass
@@ -1,54 +0,0 @@
from __future__ import absolute_import
from __future__ import division
from __future__ import print_function
import vectorized.vectorize_core as core
class Filter(object):
def _after_reset(self, observation):
return observation
def _after_step(self, observation, reward, done, info):
return observation, reward, done, info
class VectorizeFilter(core.Wrapper):
"""Vectorizes a Filter written for the non-vectorized case."""
autovectorize = False
metadata = {
'configure.required': True
}
def __init__(self, env, filter_factory, *args, **kwargs):
super(VectorizeFilter, self).__init__(env)
self.filter_factory = filter_factory
self._args = args
self._kwargs = kwargs
def _configure(self, **kwargs):
super(VectorizeFilter, self)._configure(**kwargs)
self.filter_n = [self.filter_factory(*self._args, **self._kwargs) for _ in range(self.n)]
def _reset(self):
observation_n = self.env.reset()
observation_n = [filter._after_reset(observation) for filter, observation in zip(self.filter_n, observation_n)]
return observation_n
def _step(self, action_n):
o_n, r_n, d_n, i = self.env.step(action_n)
observation_n = []
reward_n = []
done_n = []
info = i.copy()
info['n'] = []
for filter, observation, reward, done, info_i in zip(self.filter_n, o_n, r_n, d_n, i['n']):
observation, reward, done, info_i = filter._after_step(observation, reward, done, info_i)
observation_n.append(observation)
reward_n.append(reward)
done_n.append(done)
info['n'].append(info_i)
return observation_n, reward_n, done_n, info
def __str__(self):
return '<{}[{}]{}>'.format(type(self).__name__, self.filter_factory, self.env)
-93
View File
@@ -1,93 +0,0 @@
from __future__ import absolute_import
from __future__ import division
from __future__ import print_function
import gym
import weakref
import vectorized.vectorize_core as core
class Vectorize(gym.Wrapper):
"""
Given an unvectorized environment (where, e.g., the output of .step() is an observation
rather than a list of observations), turn it into a vectorized environment with a batch of size
1.
"""
metadata = {'runtime.vectorized': True}
def __init__(self, env):
super(Vectorize, self).__init__(env)
assert not env.metadata.get('runtime.vectorized')
assert self.metadata.get('runtime.vectorized')
self.n = 1
def _reset(self):
observation = self.env.reset()
return [observation]
def _step(self, action):
observation, reward, done, info = self.env.step(action[0])
return [observation], [reward], [done], {'n': [info]}
def _seed(self, seed):
return [self.env.seed(seed[0])]
class Unvectorize(core.Wrapper):
"""
Take a vectorized environment with a batch of size 1 and turn it into an unvectorized environment.
"""
autovectorize = False
metadata = {'runtime.vectorized': False}
def _configure(self, **kwargs):
super(Unvectorize, self)._configure(**kwargs)
if self.n != 1:
raise Exception('Can only disable vectorization with n=1, not n={}'.format(self.n))
def _reset(self):
observation_n = self.env.reset()
return observation_n[0]
def _step(self, action):
action_n = [action]
observation_n, reward_n, done_n, info = self.env.step(action_n)
return observation_n[0], reward_n[0], done_n[0], info['n'][0]
def _seed(self, seed):
return self.env.seed([seed])[0]
class WeakUnvectorize(Unvectorize):
def __init__(self, env, i):
self._env_ref = weakref.ref(env)
super(WeakUnvectorize, self).__init__(env)
# WeakUnvectorize won't get configure called on it
self.i = i
def _check_for_duplicate_wrappers(self):
pass # Disable this check because we need to wrap vectorized envs in multiple unvectorize wrappers
@property
def env(self):
# Called upon instantiation
if not hasattr(self, '_env_ref'):
return
env = self._env_ref()
if env is None:
raise Exception("env has been garbage collected. To keep using WeakUnvectorize, you must keep around a reference to the env object. (HINT: try assigning the env to a variable in your code.)")
return env
@env.setter
def env(self, value):
# We'll maintain our own weakref, thank you very much.
pass
def _seed(self, seed):
# We handle the seeding ourselves in the vectorized Monitor
return [seed]
def close(self):
# Don't want to close through this wrapper
pass