Optimize replay buffer as baselines

This commit is contained in:
Shangtong Zhang
2018-05-17 11:28:33 -06:00
parent 85760246e1
commit 4699748624
4 changed files with 42 additions and 73 deletions
+26 -15
View File
@@ -6,6 +6,7 @@ from gym import spaces
from gym.spaces import Box
import cv2
cv2.ocl.setUseOpenCL(False)
from collections import deque
class NoopResetEnv(gym.Wrapper):
def __init__(self, env, noop_max=30):
@@ -169,7 +170,7 @@ class LazyFrames(object):
def _force(self):
if self._out is None:
self._out = np.concatenate(self._frames, axis=2)
self._out = np.concatenate(self._frames, axis=0)
self._frames = None
return self._out
@@ -186,23 +187,33 @@ class LazyFrames(object):
return self._force()[i]
class StackFrame(gym.Wrapper):
def __init__(self, env=None, history_length=1):
super(StackFrame, self).__init__(env)
self.history_length = history_length
self.buffer = None
def __init__(self, env, k):
"""Stack k last frames.
Returns lazy array, which is much more memory efficient.
See Also
--------
baselines.common.atari_wrappers.LazyFrames
"""
gym.Wrapper.__init__(self, env)
self.k = k
self.frames = deque([], maxlen=k)
shp = env.observation_space.shape
self.observation_space = spaces.Box(low=0, high=255, shape=(shp[0] * k, shp[1], shp[2]), dtype=np.uint8)
def reset(self):
state = self.env.reset()
self.buffer = [state] * self.history_length
# return LazyFrames(self.buffer)
return np.asarray(np.vstack(self.buffer))
ob = self.env.reset()
for _ in range(self.k):
self.frames.append(ob)
return self._get_ob()
def step(self, action):
state, reward, done, info = self.env.step(action)
self.buffer.pop(0)
self.buffer.append(state)
# return LazyFrames(self.buffer), reward, done, info
return np.asarray(np.vstack(self.buffer)), reward, done, info
ob, reward, done, info = self.env.step(action)
self.frames.append(ob)
return self._get_ob(), reward, done, info
def _get_ob(self):
assert len(self.frames) == self.k
return LazyFrames(list(self.frames))
class WrapPyTorch(gym.ObservationWrapper):
# from https://github.com/ikostrikov/pytorch-a2c-ppo-acktr/blob/master/envs.py
@@ -250,7 +261,7 @@ def make_atari(env_id, frame_skip=4):
env = MaxAndSkipEnv(env, skip=4)
return env
def wrap_deepmind(env, episode_life=True, history_length=0):
def wrap_deepmind(env, episode_life=True, history_length=1):
"""Configure environment for DeepMind-style Atari.
"""
if episode_life:
+13 -57
View File
@@ -10,26 +10,15 @@ class Replay:
def __init__(self, memory_size, batch_size):
self.memory_size = memory_size
self.batch_size = batch_size
self.data = None
self.data = []
self.pos = 0
self.full = False
def feed(self, experience):
if self.data is None:
self.data = []
for unit in experience:
if np.isscalar(unit):
self.data.append(np.zeros(self.memory_size, dtype=type(unit)))
else:
self.data.append(np.zeros((self.memory_size, ) + unit.shape, unit.dtype))
for buffer_unit, exp_unit in zip(self.data, experience):
buffer_unit[self.pos] = exp_unit
self.pos += 1
if self.pos == self.memory_size:
self.full = True
self.pos = 0
if self.pos >= len(self.data):
self.data.append(experience)
else:
self.data[self.pos] = experience
self.pos = (self.pos + 1) % self.memory_size
def feed_batch(self, experience):
experience = zip(*experience)
@@ -39,47 +28,14 @@ class Replay:
def sample(self, batch_size=None):
if batch_size is None:
batch_size = self.batch_size
upper_bound = self.memory_size if self.full else self.pos
sampled_indices = np.random.randint(0, upper_bound, size=batch_size)
return [unit[sampled_indices] for unit in self.data]
sampled_indices = [np.random.randint(0, len(self.data)) for _ in range(batch_size)]
sampled_data = [self.data[ind] for ind in sampled_indices]
batch_data = list(map(lambda x: np.asarray(x), zip(*sampled_data)))
return batch_data
def size(self):
if self.full:
return self.memory_size
return self.pos
return len(self.data)
def empty(self):
return not self.full and not self.pos
class SkewedReplay:
def __init__(self, memory_size, batch_size):
memory_size = memory_size / 2
self.non_zero_reward = Replay(memory_size, batch_size / 2)
self.zero_reward = Replay(memory_size, batch_size / 2)
self.batch_size = batch_size
def feed(self, experiences):
experiences = zip(*experiences)
for exp in experiences:
if np.abs(exp[2]) < 1e-5:
self.zero_reward.feed(exp)
else:
self.non_zero_reward.feed(exp)
def sample(self):
if self.zero_reward.empty():
batch = self.non_zero_reward.sample(self.batch_size)
elif self.non_zero_reward.empty():
batch = self.zero_reward.sample(self.batch_size)
else:
non_zero_batch_size = min(self.non_zero_reward.size(), self.batch_size / 2)
zero_batch_size = min(self.zero_reward.size(), self.batch_size / 2)
batch1 = self.zero_reward.sample(zero_batch_size)
batch2 = self.non_zero_reward.sample(non_zero_batch_size)
batch = list(map(lambda seq: np.concatenate([np.asarray(x) for x in seq], axis=0), zip(batch1, batch2)))
batch = list(map(lambda x: np.asarray(x), batch))
return batch
return not len(self.data)
+2
View File
@@ -81,6 +81,8 @@ class RescaleNormalizer(BaseNormalizer):
self.coef = coef
def __call__(self, x):
if not np.isscalar(x):
x = np.asarray(x)
return self.coef * x
class ImageNormalizer(RescaleNormalizer):
+1 -1
View File
@@ -399,7 +399,7 @@ if __name__ == '__main__':
set_one_thread()
# dqn_cart_pole()
a2c_cart_pole()
# a2c_cart_pole()
# categorical_dqn_cart_pole()
# quantile_regression_dqn_cart_pole()
# n_step_dqn_cart_pole()