Optimize replay buffer

This commit is contained in:
Shangtong Zhang
2017-05-21 14:37:45 -06:00
parent ad8c8ad853
commit be9c5cdcdb
3 changed files with 59 additions and 45 deletions
+21 -3
View File
@@ -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):
+3 -4
View File
@@ -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')
+35 -38
View File
@@ -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]