From aa4ec06f2e3cf030a9b9463630109e268fdfd548 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Mon, 12 Mar 2018 22:42:47 -0600 Subject: [PATCH] N-step DQN --- agent/NStepDQN_agent.py | 81 +++++++++++++++++++++++++++++++++++++++++ agent/__init__.py | 3 +- main.py | 37 ++++++++++++++++++- 3 files changed, 119 insertions(+), 2 deletions(-) create mode 100644 agent/NStepDQN_agent.py diff --git a/agent/NStepDQN_agent.py b/agent/NStepDQN_agent.py new file mode 100644 index 0000000..f2bf05f --- /dev/null +++ b/agent/NStepDQN_agent.py @@ -0,0 +1,81 @@ +####################################################################### +# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) # +# Permission given to modify the code as long as you keep this # +# declaration at the top # +####################################################################### + +from network import * +from component import * +from utils import * +import numpy as np +import time +import os +import pickle +import torch + +class NStepDQNAgent: + def __init__(self, config): + self.config = config + self.learning_network = config.network_fn() + self.target_network = config.network_fn() + self.optimizer = config.optimizer_fn(self.learning_network.parameters()) + self.target_network.load_state_dict(self.learning_network.state_dict()) + self.task = config.task_fn() + self.policy = config.policy_fn() + + self.total_steps = 0 + self.states = self.task.reset() + self.episode_rewards = np.zeros(config.num_workers) + self.last_episode_rewards = np.zeros(config.num_workers) + + def close(self): + self.task.close() + + def save(self, file_name): + with open(file_name, 'wb') as f: + torch.save(self.learning_network.state_dict(), f) + + def iteration(self): + config = self.config + rollout = [] + states = self.states + for i in range(config.rollout_length): + q = self.learning_network.predict(states) + actions = [self.policy.sample(v) for v in q.data.cpu().numpy()] + actions = config.action_shift_fn(actions) + next_states, rewards, terminals, _ = self.task.step(actions) + self.episode_rewards += rewards + rewards = config.reward_shift_fn(rewards) + for i, terminal in enumerate(terminals): + if terminals[i]: + next_states[i] = self.task.reset(i) + self.last_episode_rewards[i] = self.episode_rewards[i] + self.episode_rewards[i] = 0 + + rollout.append([q, actions, rewards, 1 - terminals]) + states = next_states + + self.policy.update_epsilon() + self.total_steps += config.num_workers + if self.total_steps / config.num_workers % config.target_network_update_freq == 0: + self.target_network.load_state_dict(self.learning_network.state_dict()) + + self.states = states + + processed_rollout = [None] * (len(rollout)) + returns = self.target_network.predict(states).data + returns, _ = torch.max(returns, dim=1, keepdim=True) + for i in reversed(range(len(rollout))): + q, actions, rewards, terminals = rollout[i] + actions = self.learning_network.tensor(actions, torch.LongTensor).unsqueeze(1) + q = q.gather(1, Variable(actions)) + terminals = self.learning_network.tensor(terminals).unsqueeze(1) + rewards = self.learning_network.tensor(rewards).unsqueeze(1) + returns = rewards + config.discount * terminals * returns + processed_rollout[i] = [q, returns] + + q, returns= map(lambda x: torch.cat(x, dim=0), zip(*processed_rollout)) + loss = 0.5 * (q - Variable(returns)).pow(2).mean() + self.optimizer.zero_grad() + loss.backward() + self.optimizer.step() \ No newline at end of file diff --git a/agent/__init__.py b/agent/__init__.py index af6c5c1..05ce5e4 100644 --- a/agent/__init__.py +++ b/agent/__init__.py @@ -2,4 +2,5 @@ from .async_agent import * from .DQN_agent import * from .DDPG_agent import * from .A2C_agent import * -from .CategoricalDQN_agent import * \ No newline at end of file +from .CategoricalDQN_agent import * +from .NStepDQN_agent import * \ No newline at end of file diff --git a/main.py b/main.py index 3fa2517..13be964 100644 --- a/main.py +++ b/main.py @@ -378,6 +378,39 @@ def categorical_dqn_pixel_atari(name): config.categorical_n_atoms = 51 run_episodes(CategoricalDQNAgent(config)) +def n_step_dqn_cart_pole(): + config = Config() + task_fn = lambda: ClassicalControl('CartPole-v0', max_steps=200) + task = task_fn() + config.num_workers = 5 + config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers) + config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) + config.network_fn = lambda: FCNet([task.state_dim, 50, 200, task.action_dim]) + config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1) + config.discount = 0.99 + config.target_network_update_freq = 200 + config.rollout_length = 20 + config.logger = Logger('./log', logger) + run_iterations(NStepDQNAgent(config)) + +def n_step_dqn_pixel_atari(name): + config = Config() + config.history_length = 4 + task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=True, + history_length=config.history_length) + task = task_fn() + config.num_workers = 8 + config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers) + config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) + config.network_fn = lambda: NatureConvNet(config.history_length, task.action_dim, gpu=0) + config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) + config.reward_shift_fn = lambda r: np.sign(r) + config.discount = 0.99 + config.target_network_update_freq = 10000 + config.rollout_length = 20 + config.logger = Logger('./log', logger) + run_iterations(NStepDQNAgent(config)) + if __name__ == '__main__': mkdir('data') mkdir('data/video') @@ -395,9 +428,11 @@ if __name__ == '__main__': # p3o_continuous() # d3pg_continuous() # ddpg_continuous() + # n_step_dqn_cart_pole() # dqn_pixel_atari('PongNoFrameskip-v4') - categorical_dqn_pixel_atari('PongNoFrameskip-v4') + # categorical_dqn_pixel_atari('PongNoFrameskip-v4') + n_step_dqn_pixel_atari('PongNoFrameskip-v4') # async_pixel_atari('PongNoFrameskip-v4') # a3c_pixel_atari('PongNoFrameskip-v4') # a2c_pixel_atari('PongNoFrameskip-v4')