diff --git a/agent/QuantileRegressionDQN_agent.py b/agent/QuantileRegressionDQN_agent.py new file mode 100644 index 0000000..56ba90f --- /dev/null +++ b/agent/QuantileRegressionDQN_agent.py @@ -0,0 +1,101 @@ +####################################################################### +# 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 QuantileRegressionDQNAgent: + 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.criterion = nn.MSELoss() + self.target_network.load_state_dict(self.learning_network.state_dict()) + self.task = config.task_fn() + self.replay = config.replay_fn() + self.policy = config.policy_fn() + self.total_steps = 0 + self.quantile_weight = 1.0 / self.config.num_quantiles + self.cumulative_density = self.learning_network.tensor( + (2 * np.arange(self.config.num_quantiles) + 1) / (2.0 * self.config.num_quantiles)) + + def huber(self, x): + cond = (x < 1.0).float().detach() + return 0.5 * x.pow(2) * cond + (x.abs() - 0.5) * (1 - cond) + + def episode(self, deterministic=False): + episode_start_time = time.time() + state = self.task.reset() + total_reward = 0.0 + steps = 0 + while True: + value = self.learning_network.predict(np.stack([self.task.normalize_state(state)])).squeeze(0).data + value = (value * self.quantile_weight).sum(-1).cpu().numpy().flatten() + if deterministic: + action = np.argmax(value) + elif self.total_steps < self.config.exploration_steps: + action = np.random.randint(0, len(value)) + else: + action = self.policy.sample(value) + next_state, reward, done, _ = self.task.step(action) + total_reward += reward + reward = self.config.reward_shift_fn(reward) + if not deterministic: + self.replay.feed([state, action, reward, next_state, int(done)]) + self.total_steps += 1 + steps += 1 + state = next_state + if done: + break + if not deterministic and self.total_steps > self.config.exploration_steps: + experiences = self.replay.sample() + states, actions, rewards, next_states, terminals = experiences + states = self.task.normalize_state(states) + next_states = self.task.normalize_state(next_states) + + quantiles_next = self.target_network.predict(next_states).data + q_next = (quantiles_next * self.quantile_weight).sum(-1) + _, a_next = torch.max(q_next, dim=1) + a_next = a_next.view(-1, 1, 1).expand(-1, -1, quantiles_next.size(2)) + quantiles_next = quantiles_next.gather(1, a_next).squeeze(1) + + rewards = self.learning_network.tensor(rewards) + terminals = self.learning_network.tensor(terminals) + quantiles_next = rewards.view(-1, 1) + self.config.discount * (1 - terminals.view(-1, 1)) * quantiles_next + + quantiles = self.learning_network.predict(states) + actions = self.learning_network.tensor(actions, torch.LongTensor) + actions = actions.view(-1, 1, 1).expand(-1, -1, quantiles.size(2)) + quantiles = quantiles.gather(1, Variable(actions)).squeeze(1) + + diff = Variable(quantiles_next) - quantiles + loss = self.huber(diff) * Variable(self.cumulative_density.view(1, -1) - (diff.data < 0).float()).abs() + + self.optimizer.zero_grad() + loss.sum(-1).mean().backward() + self.optimizer.step() + if not deterministic and self.total_steps % self.config.target_network_update_freq == 0: + self.target_network.load_state_dict(self.learning_network.state_dict()) + if not deterministic and self.total_steps > self.config.exploration_steps: + self.policy.update_epsilon() + episode_time = time.time() - episode_start_time + self.config.logger.debug('episode steps %d, episode time %f, time per step %f' % + (steps, episode_time, episode_time / float(steps))) + return total_reward, steps + + def save(self, file_name): + with open(file_name, 'wb') as f: + torch.save(self.learning_network.state_dict(), f) + + def close(self): + pass diff --git a/agent/__init__.py b/agent/__init__.py index 05ce5e4..7d48e63 100644 --- a/agent/__init__.py +++ b/agent/__init__.py @@ -3,4 +3,5 @@ from .DQN_agent import * from .DDPG_agent import * from .A2C_agent import * from .CategoricalDQN_agent import * -from .NStepDQN_agent import * \ No newline at end of file +from .NStepDQN_agent import * +from .QuantileRegressionDQN_agent import * \ No newline at end of file diff --git a/main.py b/main.py index 01c5e21..8344db9 100644 --- a/main.py +++ b/main.py @@ -411,6 +411,24 @@ def n_step_dqn_pixel_atari(name): config.logger = Logger('./log', logger) run_iterations(NStepDQNAgent(config)) +def quantile_regression_dqn_cart_pole(): + config = Config() + config.task_fn = lambda: ClassicalControl('CartPole-v0', max_steps=200) + task = config.task_fn() + config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001) + config.network_fn = lambda: QuantileFCNet(task.state_dim, task.action_dim, config.num_quantiles) + config.policy_fn = lambda: GreedyPolicy(epsilon=0.1, final_step=10000, min_epsilon=0.1) + config.replay_fn = lambda: Replay(memory_size=10000, batch_size=10) + config.discount = 0.99 + config.target_network_update_freq = 200 + config.exploration_steps = 100 + config.logger = Logger('./log', logger, skip=True) + # config.logger = Logger('./log', logger) + config.test_interval = 100 + config.test_repetitions = 50 + config.num_quantiles = 20 + run_episodes(QuantileRegressionDQNAgent(config)) + if __name__ == '__main__': mkdir('data') mkdir('data/video') @@ -421,9 +439,10 @@ if __name__ == '__main__': # dqn_cart_pole() # categorical_dqn_cart_pole() + quantile_regression_dqn_cart_pole() # async_cart_pole() # a3c_cart_pole() - a2c_cart_pole() + # a2c_cart_pole() # a3c_continuous() # p3o_continuous() # d3pg_continuous() diff --git a/network/base_network.py b/network/base_network.py index 1fee1d0..3c3e143 100644 --- a/network/base_network.py +++ b/network/base_network.py @@ -101,3 +101,8 @@ class CategoricalNet(BasicNet): if to_numpy: return prob.cpu().data.numpy() return prob + +class QuantileNet(BasicNet): + def predict(self, x, to_numpy=False): + quantiles = self.forward(x) + return quantiles.view((-1, self.n_actions, self.n_quantiles)) diff --git a/network/shallow_network.py b/network/shallow_network.py index 1094a81..65dbb8e 100644 --- a/network/shallow_network.py +++ b/network/shallow_network.py @@ -73,3 +73,22 @@ class CategoricalFCNet(nn.Module, CategoricalNet): phi = F.relu(self.fc1(x)) phi = F.relu(self.fc2(phi)) return phi + +class QuantileFCNet(nn.Module, QuantileNet): + def __init__(self, state_dim, n_actions, n_quantiles, gpu=0): + super(QuantileFCNet, self).__init__() + self.n_actions = n_actions + self.n_quantiles = n_quantiles + + hidden_size = 64 + self.fc1 = nn.Linear(state_dim, hidden_size) + self.fc2 = nn.Linear(hidden_size, hidden_size) + self.fc3 = nn.Linear(hidden_size, n_actions * n_quantiles) + BasicNet.__init__(self, gpu) + + def forward(self, x): + x = self.variable(x) + phi = F.relu(self.fc1(x)) + phi = F.relu(self.fc2(phi)) + quantiles = self.fc3(phi) + return quantiles \ No newline at end of file diff --git a/utils/config.py b/utils/config.py index 4643e5c..169597c 100644 --- a/utils/config.py +++ b/utils/config.py @@ -58,3 +58,4 @@ class Config: self.categorical_v_min = -10 self.categorical_v_max = 10 self.categorical_n_atoms = 51 + self.num_quantiles = 10