From 469974862440c1e3f6bd5e8c1d08a8b7eca0a6de Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Thu, 17 May 2018 11:28:33 -0600 Subject: [PATCH] Optimize replay buffer as baselines --- deep_rl/component/atari_wrapper.py | 41 ++++++++++------- deep_rl/component/replay.py | 70 ++++++------------------------ deep_rl/utils/normalizer.py | 2 + examples.py | 2 +- 4 files changed, 42 insertions(+), 73 deletions(-) diff --git a/deep_rl/component/atari_wrapper.py b/deep_rl/component/atari_wrapper.py index a4df712..9e83a7b 100644 --- a/deep_rl/component/atari_wrapper.py +++ b/deep_rl/component/atari_wrapper.py @@ -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: diff --git a/deep_rl/component/replay.py b/deep_rl/component/replay.py index a9bef35..47b75f6 100644 --- a/deep_rl/component/replay.py +++ b/deep_rl/component/replay.py @@ -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) diff --git a/deep_rl/utils/normalizer.py b/deep_rl/utils/normalizer.py index 323db25..3bda240 100644 --- a/deep_rl/utils/normalizer.py +++ b/deep_rl/utils/normalizer.py @@ -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): diff --git a/examples.py b/examples.py index 1e37af0..b4fb7be 100644 --- a/examples.py +++ b/examples.py @@ -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()