diff --git a/agent/CategoricalDQN_agent.py b/agent/CategoricalDQN_agent.py new file mode 100644 index 0000000..d2eaccc --- /dev/null +++ b/agent/CategoricalDQN_agent.py @@ -0,0 +1,105 @@ +####################################################################### +# 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 CategoricalDQNAgent: + 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.atoms = self.learning_network.tensor( + np.linspace(config.categorical_v_min, + config.categorical_v_max, + config.categorical_n_atoms)) + self.delta_atom = (config.categorical_v_max - config.categorical_v_min) / float(config.categorical_n_atoms - 1) + + 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 = torch.mm(value, self.atoms.unsqueeze(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 += np.sum(reward * self.config.reward_weight) + 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) + prob_next = self.target_network.predict(next_states).data + q_next = (prob_next * self.atoms).sum(-1) + _, a_next = torch.max(q_next, dim=1) + a_next = a_next.view(-1, 1, 1).expand(-1, -1, prob_next.size(2)) + prob_next = prob_next.gather(1, a_next).squeeze(1) + + rewards = self.learning_network.tensor(rewards) + atoms_next = rewards.view(-1, 1) + self.config.discount * self.atoms.view(1, -1) + epsilon = 1e-5 + atoms_next.clamp_(self.config.categorical_v_min + epsilon, self.config.categorical_v_max - epsilon) + b = (atoms_next - self.config.categorical_v_min) / self.delta_atom + l = b.floor() + u = b.ceil() + d_m_l = (u - b) * prob_next + d_m_u = (b - l) * prob_next + target_prob = self.learning_network.tensor(np.zeros(prob_next.size())) + for i in range(target_prob.size(0)): + target_prob[i].index_add_(0, l[i].long(), d_m_l[i]) + target_prob[i].index_add_(0, u[i].long(), d_m_u[i]) + + prob = self.learning_network.predict(states) + actions = self.learning_network.tensor(actions, torch.LongTensor) + actions = actions.view(-1, 1, 1).expand(-1, -1, prob.size(2)) + prob = prob.gather(1, Variable(actions)).squeeze(1) + loss = -(Variable(target_prob) * prob.log()).sum(-1).mean() + self.optimizer.zero_grad() + loss.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 b4a371e..af6c5c1 100644 --- a/agent/__init__.py +++ b/agent/__init__.py @@ -1,4 +1,5 @@ from .async_agent import * from .DQN_agent import * from .DDPG_agent import * -from .A2C_agent import * \ No newline at end of file +from .A2C_agent import * +from .CategoricalDQN_agent import * \ No newline at end of file diff --git a/main.py b/main.py index 293a4dd..cc2485c 100644 --- a/main.py +++ b/main.py @@ -335,6 +335,25 @@ def ddpg_continuous(): config.logger = Logger('./log', logger) run_episodes(DDPGAgent(config)) +def categorical_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: CategoricalFCNet(task.state_dim, task.action_dim, config.categorical_n_atoms) + 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 = 0 + config.logger = Logger('./log', logger) + config.test_interval = 100 + config.test_repetitions = 50 + config.categorical_v_max = 200 + config.categorical_v_min = -config.categorical_v_min + config.categorical_n_atoms = 10 + run_episodes(CategoricalDQNAgent(config)) + if __name__ == '__main__': mkdir('data') mkdir('data/video') @@ -343,7 +362,8 @@ if __name__ == '__main__': # logger.setLevel(logging.DEBUG) logger.setLevel(logging.INFO) - dqn_cart_pole() + # dqn_cart_pole() + categorical_dqn_cart_pole() # async_cart_pole() # a3c_cart_pole() # a2c_cart_pole() @@ -361,7 +381,7 @@ if __name__ == '__main__': # async_pixel_atari('BreakoutNoFrameskip-v4') # a3c_pixel_atari('BreakoutNoFrameskip-v4') - dqn_ram_atari('Pong-ramNoFrameskip-v4') + # dqn_ram_atari('Pong-ramNoFrameskip-v4') # acvp.train('PongNoFrameskip-v4') diff --git a/network/base_network.py b/network/base_network.py index 52e0501..82abbb0 100644 --- a/network/base_network.py +++ b/network/base_network.py @@ -91,4 +91,13 @@ class DuelingNet(BasicNet): q = value.expand_as(advantange) + (advantange - advantange.mean(1).expand_as(advantange)) if to_numpy: return q.cpu().data.numpy() - return q \ No newline at end of file + return q + +class CategoricalNet(BasicNet): + def predict(self, x, to_numpy=False): + phi = self.forward(x) + pre_prob = self.fc_categorical(phi).view((-1, self.n_actions, self.n_atoms)) + prob = F.softmax(pre_prob, dim=-1) + if to_numpy: + return pre_prob.cpu().data.numpy() + return prob diff --git a/network/shallow_network.py b/network/shallow_network.py index 83f3015..1094a81 100644 --- a/network/shallow_network.py +++ b/network/shallow_network.py @@ -55,3 +55,21 @@ class ActorCriticFCNet(nn.Module, ActorCriticNet): x = F.relu(self.fc1(x)) phi = self.fc2(x) return phi + +class CategoricalFCNet(nn.Module, CategoricalNet): + def __init__(self, state_dim, n_actions, n_atoms, gpu=0): + super(CategoricalFCNet, self).__init__() + self.n_actions = n_actions + self.n_atoms = n_atoms + + hidden_size = 64 + self.fc1 = nn.Linear(state_dim, hidden_size) + self.fc2 = nn.Linear(hidden_size, hidden_size) + self.fc_categorical = nn.Linear(hidden_size, n_actions * n_atoms) + BasicNet.__init__(self, gpu) + + def forward(self, x): + x = self.variable(x) + phi = F.relu(self.fc1(x)) + phi = F.relu(self.fc2(phi)) + return phi diff --git a/utils/config.py b/utils/config.py index 689475c..4643e5c 100644 --- a/utils/config.py +++ b/utils/config.py @@ -55,3 +55,6 @@ class Config: self.rollout_length = None self.value_loss_weight = 1.0 self.iteration_log_interval = 30 + self.categorical_v_min = -10 + self.categorical_v_max = 10 + self.categorical_n_atoms = 51