mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Optimize replay buffer as baselines
This commit is contained in:
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user