From be9c5cdcdb84ce314633871ed339a18de28c5dca Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Sun, 21 May 2017 14:37:45 -0600 Subject: [PATCH] Optimize replay buffer --- dqn_agent.py | 24 ++++++++++++++--- main.py | 7 +++-- replay.py | 73 +++++++++++++++++++++++++--------------------------- 3 files changed, 59 insertions(+), 45 deletions(-) diff --git a/dqn_agent.py b/dqn_agent.py index d5c9ee8..06ed10b 100644 --- a/dqn_agent.py +++ b/dqn_agent.py @@ -8,6 +8,9 @@ from network import * from replay import * from policy import * import numpy as np +import time +import psutil +import os class DQNAgent: def __init__(self, @@ -35,6 +38,7 @@ class DQNAgent: self.explore_steps = explore_steps self.history_length = history_length self.logger = logger + self.process = psutil.Process(os.getpid()) def get_state(self, history_buffer): if self.history_length > 1: @@ -42,6 +46,7 @@ class DQNAgent: return history_buffer[0] def episode(self): + episode_start_time = time.time() state = self.task.reset() history_buffer = [state] * self.history_length total_reward = 0.0 @@ -58,25 +63,38 @@ class DQNAgent: self.replay.feed([state, action, reward, next_state, int(done)]) steps += 1 self.total_steps += 1 - self.logger.debug('steps %d, reward %f, action %d' % (steps, reward, action)) if done: break if self.total_steps > self.explore_steps: + sample_start_time = time.time() experiences = self.replay.sample() + self.logger.debug('sample time %f' % (time.time() - sample_start_time)) states, actions, rewards, next_states, terminals = experiences + predict_start_time = time.time() targets = self.learning_network.predict(states) q_next = self.target_network.predict(next_states) + self.logger.debug('prediction time %f' % (time.time() - predict_start_time)) q_next = np.max(q_next, axis=1) q_next = np.where(terminals, 0, q_next) q_next = rewards + self.discount * q_next targets[np.arange(len(actions)), actions] = q_next - self.logger.debug('start minibatch') + minibatch_start_time = time.time() self.learning_network.learn(states, targets) - self.logger.debug('minibatch ended') + self.logger.debug('minibatch time %f' % (time.time() - minibatch_start_time)) if self.total_steps % self.target_network_update_freq == 0: self.target_network.load_state_dict(self.learning_network.state_dict()) if self.total_steps > self.explore_steps: self.policy.update_epsilon() + episode_time = time.time() - episode_start_time + info = self.process.memory_full_info() + if hasattr(info, 'swap'): + info_stat = info.swap + elif hasattr(info, 'pfaults'): + info_stat = info.pfaults + else: + info_stat = -1 + self.logger.debug('episode steps %d, episode time %f, time per step %f, memory_info %d' % + (steps, episode_time, episode_time / float(steps), info_stat)) return total_reward def run(self): diff --git a/main.py b/main.py index ff8a896..214c8eb 100644 --- a/main.py +++ b/main.py @@ -77,7 +77,7 @@ def dqn_pixel_atari(name): config['optimizer_fn'] = lambda params: torch.optim.RMSprop(params, lr=0.00025) config['network_fn'] = lambda optimizer_fn: ConvNet(4, 6, optimizer_fn) config['policy_fn'] = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) - config['replay_fn'] = lambda: Replay(memory_size=1000000, batch_size=32) + config['replay_fn'] = lambda: Replay(memory_size=200000, batch_size=32) config['discount'] = 0.99 config['target_network_update_freq'] = 10000 config['step_limit'] = 0 @@ -88,11 +88,10 @@ def dqn_pixel_atari(name): agent.run() if __name__ == '__main__': - # gym.logger.setLevel(logging.DEBUG) - gym.logger.setLevel(logging.INFO) + gym.logger.setLevel(logging.DEBUG) + # gym.logger.setLevel(logging.INFO) # async_cart_pole() # async_lunar_lander() # dqn_cart_pole() - # dqn_mountain_car() # actor_critic_cart_pole() dqn_pixel_atari('Breakout-v0') diff --git a/replay.py b/replay.py index 1e7cb89..6dfafb2 100644 --- a/replay.py +++ b/replay.py @@ -11,47 +11,44 @@ class Replay: self.memory_size = memory_size self.batch_size = batch_size - self.states = [] - self.actions = [] - self.rewards = [] - self.next_states = [] - self.terminals = [] + self.states = None + self.actions = np.empty(self.memory_size, dtype=np.int8) + 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 - self.states.append(state) - self.actions.append(action) - self.rewards.append(reward) - self.next_states.append(next_state) - self.terminals.append(done) - if len(self.terminals) > self.memory_size: - self.states.pop(0) - self.actions.pop(0) - self.rewards.pop(0) - self.next_states.pop(0) - self.terminals.pop(0) + + if self.states is None: + self.states = np.empty((self.memory_size, ) + state.shape) + self.next_states = np.empty((self.memory_size, ) + state.shape) + + 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): - if len(self.terminals) >= self.batch_size: - sampled_indices = np.arange(len(self.terminals)) - np.random.shuffle(sampled_indices) - sampled_indices = sampled_indices[: self.batch_size] - sampled_states = [] - sampled_actions = [] - sampled_rewards = [] - sampled_next_states = [] - sampled_terminals = [] - for ind in sampled_indices: - sampled_states.append(self.states[ind]) - sampled_actions.append(self.actions[ind]) - sampled_rewards.append(self.rewards[ind]) - sampled_next_states.append(self.next_states[ind]) - sampled_terminals.append(self.terminals[ind]) - return [np.asarray(sampled_states), - np.asarray(sampled_actions), - np.asarray(sampled_rewards), - np.asarray(sampled_next_states), - np.asarray(sampled_terminals)] - return None - + upper_bound = self.memory_size if self.full else self.pos + sampled_indices = np.random.randint(0, upper_bound, size=self.batch_size) + sampled_states = self.states[sampled_indices] + sampled_actions = self.actions[sampled_indices] + sampled_rewards = self.rewards[sampled_indices] + sampled_next_states = self.next_states[sampled_indices] + sampled_terminals = self.terminals[sampled_indices] + return [sampled_states, + sampled_actions, + sampled_rewards, + sampled_next_states, + sampled_terminals]