From 92f03661759885001fb5f6ebdcdfa5944c0d2411 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Sat, 21 Apr 2018 11:21:48 -0600 Subject: [PATCH] Refactor replay --- component/replay.py | 264 ++++++-------------------------------------- main.py | 14 +-- 2 files changed, 41 insertions(+), 237 deletions(-) diff --git a/component/replay.py b/component/replay.py index 18225b7..fda18a1 100644 --- a/component/replay.py +++ b/component/replay.py @@ -10,260 +10,64 @@ import random import torch.multiprocessing as mp class Replay: - def __init__(self, memory_size, batch_size, dtype=np.float32): - self.memory_size = memory_size - self.batch_size = batch_size - self.dtype = dtype - - self.states = None - self.actions = np.empty(self.memory_size, dtype=np.uint8) - self.rewards = np.empty(self.memory_size) - self.next_states = None - self.terminals = np.empty(self.memory_size, dtype=np.uint8) - - self.pos = 0 - self.full = False - - - def feed(self, experience): - state, action, reward, next_state, done = experience - - if self.states is None: - self.states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype) - self.next_states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype) - - self.states[self.pos][:] = state - self.actions[self.pos] = action - self.rewards[self.pos] = reward - self.next_states[self.pos][:] = next_state - self.terminals[self.pos] = done - - self.pos += 1 - if self.pos == self.memory_size: - self.full = True - self.pos = 0 - - def sample(self): - upper_bound = self.memory_size if self.full else self.pos - sampled_indices = np.random.randint(0, upper_bound, size=self.batch_size) - return [self.states[sampled_indices], - self.actions[sampled_indices], - self.rewards[sampled_indices], - self.next_states[sampled_indices], - self.terminals[sampled_indices]] - -class HybridRewardReplay: - def __init__(self, memory_size, batch_size, dtype=np.float32): - self.memory_size = memory_size - self.batch_size = batch_size - self.dtype = dtype - - self.states = None - self.actions = np.empty(self.memory_size, dtype=np.uint8) - self.rewards = None - self.next_states = None - self.terminals = np.empty(self.memory_size, dtype=np.uint8) - - self.pos = 0 - self.full = False - - - def feed(self, experience): - state, action, reward, next_state, done = experience - - if self.states is None: - self.rewards = np.empty((self.memory_size, ) + reward.shape, dtype=self.dtype) - self.states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype) - self.next_states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype) - - self.states[self.pos][:] = state - self.actions[self.pos] = action - self.rewards[self.pos][:] = reward - self.next_states[self.pos][:] = next_state - self.terminals[self.pos] = done - - self.pos += 1 - if self.pos == self.memory_size: - self.full = True - self.pos = 0 - - def sample(self): - upper_bound = self.memory_size if self.full else self.pos - sampled_indices = np.random.randint(0, upper_bound, size=self.batch_size) - return [self.states[sampled_indices], - self.actions[sampled_indices], - self.rewards[sampled_indices], - self.next_states[sampled_indices], - self.terminals[sampled_indices]] - -class SharedReplay: - def __init__(self, memory_size, batch_size, state_shape, action_shape): - self.memory_size = memory_size - self.batch_size = batch_size - - self.states = torch.zeros((self.memory_size, ) + state_shape) - self.actions = torch.zeros((self.memory_size, ) + action_shape) - self.rewards = torch.zeros(self.memory_size) - self.next_states = torch.zeros((self.memory_size, ) + state_shape) - self.terminals = torch.zeros(self.memory_size) - - self.states.share_memory_() - self.actions.share_memory_() - self.rewards.share_memory_() - self.next_states.share_memory_() - self.terminals.share_memory_() - - self.pos = 0 - self.full = False - self.buffer_lock = mp.Lock() - - def feed_(self, experience): - state, action, reward, next_state, done = experience - self.states[self.pos][:] = torch.FloatTensor(state) - self.actions[self.pos][:] = torch.FloatTensor(action) - self.rewards[self.pos] = reward - self.next_states[self.pos][:] = torch.FloatTensor(next_state) - self.terminals[self.pos] = done - - self.pos += 1 - if self.pos == self.memory_size: - self.full = True - self.pos = 0 - - def size(self): - if self.full: - return self.memory_size - return self.pos - - def sample_(self): - upper_bound = self.memory_size if self.full else self.pos - sampled_indices = torch.LongTensor(np.random.randint(0, upper_bound, size=self.batch_size)) - return [self.states[sampled_indices], - self.actions[sampled_indices], - self.rewards[sampled_indices], - self.next_states[sampled_indices], - self.terminals[sampled_indices]] - - def feed(self, experience): - with self.buffer_lock: - self.feed_(experience) - - def sample(self): - with self.buffer_lock: - return self.sample_() - - def state_dict(self): - return dict((key, getattr(self, key)) for key in ['actions', 'states', 'rewards', 'next_states', 'terminals', 'pos']) - - def load_state_dict(self, state): - for key in ['actions', 'states', 'rewards', 'next_states', 'terminals', 'pos']: - val = state[key] - setattr(self, key, val) - - def save(self, file_name): - with open(file_name, 'wb') as f: - torch.save(self.state_dict(), f) - - def load(self, file_name): - state = torch.load(file_name) - self.load_state_dict(state) - -class HighDimActionReplay: - def __init__(self, memory_size, batch_size, dtype=np.float32): - self.memory_size = memory_size - self.batch_size = batch_size - self.dtype = dtype - - self.states = None - self.actions = None - self.rewards = np.empty(self.memory_size) - self.next_states = None - self.terminals = np.empty(self.memory_size, dtype=np.int8) - - self.pos = 0 - self.full = False - - - def feed(self, experience): - state, action, reward, next_state, done = experience - - if self.states is None: - self.states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype) - self.actions = np.empty((self.memory_size, ) + action.shape) - self.next_states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype) - - self.states[self.pos][:] = state - self.actions[self.pos][:] = action - self.rewards[self.pos] = reward - self.next_states[self.pos][:] = next_state - self.terminals[self.pos] = done - - self.pos += 1 - if self.pos == self.memory_size: - self.full = True - self.pos = 0 - - def size(self): - if self.full: - return self.memory_size - return self.pos - - def sample(self): - upper_bound = self.memory_size if self.full else self.pos - sampled_indices = np.random.randint(0, upper_bound, size=self.batch_size) - return [self.states[sampled_indices], - self.actions[sampled_indices], - self.rewards[sampled_indices], - self.next_states[sampled_indices], - self.terminals[sampled_indices]] - -class GeneralReplay: def __init__(self, memory_size, batch_size): - self.buffer = [] self.memory_size = memory_size self.batch_size = batch_size + self.data = None - def feed(self, experiences): - for experience in zip(*experiences): - self.feed_single(experience) + self.pos = 0 + self.full = False - def feed_single(self, experience): - self.buffer.append(experience) - if len(self.buffer) > self.memory_size: - del self.buffer[0] + 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 + + def feed_batch(self, experience): + experience = zip(*experience) + for exp in experience: + self.feed(exp) def sample(self, batch_size=None): if batch_size is None: batch_size = self.batch_size - sampled = zip(*random.sample(self.buffer, batch_size)) - return sampled - - def clear(self): - self.buffer = [] - - def full(self): - return len(self.buffer) == self.memory_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] def size(self): - return len(self.buffer) + if self.full: + return self.memory_size + return self.pos def empty(self): - return not len(self.buffer) + 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 = GeneralReplay(memory_size, batch_size / 2) - self.zero_reward = GeneralReplay(memory_size, batch_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_single(exp) + self.zero_reward.feed(exp) else: - self.non_zero_reward.feed_single(exp) + self.non_zero_reward.feed(exp) def sample(self): if self.zero_reward.empty(): diff --git a/main.py b/main.py index ffe14ba..eaeb51b 100644 --- a/main.py +++ b/main.py @@ -132,7 +132,7 @@ def dqn_pixel_atari(name): config.network_fn = lambda state_dim, action_dim: VanillaNet(action_dim, NatureConvBody(), gpu=0) # config.network_fn = lambda state_dim, action_dim: DuelingNet(action_dim, NatureConvBody(), gpu=0) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) - config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32, dtype=np.uint8) + config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32) config.state_normalizer = ImageNormalizer() config.reward_normalizer = SignNormalizer() config.discount = 0.99 @@ -173,7 +173,7 @@ def categorical_dqn_pixel_atari(name): config.network_fn = lambda state_dim, action_dim: \ CategoricalNet(action_dim, config.categorical_n_atoms, NatureConvBody(), gpu=1) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) - config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32, dtype=np.uint8) + config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32) config.discount = 0.99 config.state_normalizer = ImageNormalizer() config.reward_normalizer = SignNormalizer() @@ -195,7 +195,7 @@ def quantile_regression_dqn_pixel_atari(name): config.network_fn = lambda state_dim, action_dim: \ QuantileNet(action_dim, config.num_quantiles, NatureConvBody(), gpu=2) config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.01) - config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32, dtype=np.uint8) + config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32) config.state_normalizer = ImageNormalizer() config.reward_normalizer = SignNormalizer() config.discount = 0.99 @@ -258,7 +258,7 @@ def dqn_ram_atari(name): config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) config.network_fn = lambda state_dim, action_dim: VanillaNet(action_dim, TwoLayerFCBody(state_dim), gpu=2) config.policy_fn = lambda: GreedyPolicy(epsilon=0.1, final_step=1000000, min_epsilon=0.1) - config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32, dtype=np.uint8) + config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32) config.state_normalizer = RescaleNormalizer(1.0 / 128) config.reward_normalizer = SignNormalizer() config.discount = 0.99 @@ -310,17 +310,17 @@ def ddpg_continuous(): # config.task_fn = lambda: Pendulum(log_dir=log_dir) # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1', log_dir=log_dir) # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1', log_dir=log_dir) - config.task_fn = lambda: Roboschool('RoboschoolHopper-v1', log_dir=log_dir) + config.task_fn = lambda: Roboschool('RoboschoolHopper-v1') # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1', log_dir=log_dir) # config.task_fn = lambda: Roboschool('RoboschoolWalker2d-v1', log_dir=log_dir) # config.task_fn = lambda: DMControl('cartpole', 'balance', log_dir=log_dir) # config.task_fn = lambda: DMControl('finger', 'spin', log_dir=log_dir) - config.evaluation_env = config.task_fn() + config.evaluation_env = Roboschool('RoboschoolHopper-v1', log_dir=log_dir) config.actor_network_fn = lambda state_dim, action_dim: DeterministicActorNet(state_dim, action_dim) config.critic_network_fn = lambda state_dim, action_dim: DeterministicCriticNet(state_dim, action_dim) config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4) config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4) - config.replay_fn = lambda: HighDimActionReplay(memory_size=1000000, batch_size=64) + config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=64) config.discount = 0.99 config.state_normalizer = RunningStatsNormalizer() config.random_process_fn = \