diff --git a/agent/A2C_agent.py b/agent/A2C_agent.py index ad9826c..ff62ff8 100644 --- a/agent/A2C_agent.py +++ b/agent/A2C_agent.py @@ -46,7 +46,6 @@ class A2CAgent: 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 diff --git a/agent/NStepDQN_agent.py b/agent/NStepDQN_agent.py index 2ab81a6..5ad05c1 100644 --- a/agent/NStepDQN_agent.py +++ b/agent/NStepDQN_agent.py @@ -48,7 +48,6 @@ class NStepDQNAgent: 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 diff --git a/agent/PPO_agent.py b/agent/PPO_agent.py new file mode 100644 index 0000000..6737f60 --- /dev/null +++ b/agent/PPO_agent.py @@ -0,0 +1,110 @@ +####################################################################### +# 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 # +####################################################################### + +import numpy as np +import torch.multiprocessing as mp +from network import * +from utils import * +from component import * +import pickle +import os +import time + +class PPOAgent: + def __init__(self, config): + self.config = config + self.task = config.task_fn() + self.actor = config.actor_network_fn(self.task.state_dim, self.task.action_dim) + self.critic = config.critic_network_fn(self.task.state_dim, self.task.action_dim) + self.actor_opt = config.actor_optimizer_fn(self.actor.parameters()) + self.critic_opt = config.critic_optimizer_fn(self.critic.parameters()) + self.total_steps = 0 + self.episode_rewards = np.zeros(config.num_workers) + self.last_episode_rewards = np.zeros(config.num_workers) + self.state_normalizer = Normalizer(self.task.state_dim) + self.states = self.task.reset() + self.states = self.state_normalizer(self.states) + + def close(self): + self.task.close() + + def save(self, file_name): + pass + # with open(file_name, 'wb') as f: + # torch.save(self.network.state_dict(), f) + + def iteration(self): + config = self.config + rollout = [] + states = self.states + for i in range(config.rollout_length): + mean, std, log_std = self.actor.predict(states) + values = self.critic.predict(states) + dist = torch.distributions.Normal(mean, std) + actions = dist.sample() + log_probs = dist.log_prob(actions).detach() + log_probs = torch.sum(log_probs, dim=1, keepdim=True) + next_states, rewards, terminals, _ = self.task.step(actions.data.cpu().numpy()) + self.episode_rewards += rewards + rewards = config.reward_shift_fn(rewards) + for i, terminal in enumerate(terminals): + if terminals[i]: + self.last_episode_rewards[i] = self.episode_rewards[i] + self.episode_rewards[i] = 0 + next_states = self.state_normalizer(next_states) + rollout.append([states, values, actions, log_probs, rewards, 1 - terminals]) + states = next_states + + self.states = states + pending_value = self.critic.predict(states) + rollout.append([states, pending_value, None, None, None, None]) + + processed_rollout = [None] * (len(rollout) - 1) + advantages = self.actor.tensor(np.zeros((config.num_workers, 1))) + returns = pending_value.data + for i in reversed(range(len(rollout) - 1)): + states, value, actions, log_probs, rewards, terminals = rollout[i] + terminals = self.actor.tensor(terminals).unsqueeze(1) + rewards = self.actor.tensor(rewards).unsqueeze(1) + actions = self.actor.variable(actions) + states = self.actor.variable(states) + next_value = rollout[i + 1][1] + returns = rewards + config.discount * terminals * returns + if not config.use_gae: + advantages = returns - value.data + else: + td_error = rewards + config.discount * terminals * next_value.data - value.data + advantages = advantages * config.gae_tau * config.discount * terminals + td_error + processed_rollout[i] = [states, actions, log_probs, returns, advantages] + + states, actions, log_probs_old, returns, advantages = map(lambda x: torch.cat(x, dim=0), zip(*processed_rollout)) + advantages = (advantages - advantages.mean()) / advantages.std() + advantages = Variable(advantages) + + for k in range(config.optimization_epochs): + mean, std, log_std = self.actor.predict(states) + dist = torch.distributions.Normal(mean, std) + log_probs = dist.log_prob(actions) + log_probs = torch.sum(log_probs, dim=1, keepdim=True) + ratio = (log_probs - log_probs_old).exp() + obj = ratio * advantages + obj_clipped = ratio.clamp(1.0 - self.config.ppo_ratio_clip, 1.0 + self.config.ppo_ratio_clip) * advantages + policy_loss = -torch.min(obj, obj_clipped).mean(0) + + v = self.critic.predict(states) + value_loss = 0.5 * (Variable(returns) - v).pow(2).mean() + + self.actor_opt.zero_grad() + self.critic_opt.zero_grad() + policy_loss.backward() + value_loss.backward() + nn.utils.clip_grad_norm(self.actor.parameters(), config.gradient_clip) + nn.utils.clip_grad_norm(self.critic.parameters(), config.gradient_clip) + self.actor_opt.step() + self.critic_opt.step() + + steps = config.rollout_length * config.num_workers + self.total_steps += steps diff --git a/agent/__init__.py b/agent/__init__.py index 41f28d5..326c8bb 100644 --- a/agent/__init__.py +++ b/agent/__init__.py @@ -5,3 +5,4 @@ from .A2C_agent import * from .CategoricalDQN_agent import * from .NStepDQN_agent import * from .QuantileRegressionDQN_agent import * +from .PPO_agent import * diff --git a/component/task.py b/component/task.py index 3a8422c..8a878b0 100644 --- a/component/task.py +++ b/component/task.py @@ -151,7 +151,10 @@ def sub_task(parent_pipe, pipe, task_fn, rank, log_dir): while True: op, data = pipe.recv() if op == 'step': - pipe.send(task.step(data)) + ob, reward, done, info = task.step(data) + if done: + ob = task.reset() + pipe.send([ob, reward, done, info]) elif op == 'reset': pipe.send(task.reset()) elif op == 'exit': diff --git a/main.py b/main.py index 6b5ea0e..d1e960b 100644 --- a/main.py +++ b/main.py @@ -203,105 +203,32 @@ def dqn_ram_atari(name): # config.double_q = False run_episodes(DQNAgent(config)) -# def a3c_continuous(): -# config = Config() -# config.task_fn = lambda: Pendulum() -# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') -# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') -# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') -# task = config.task_fn() -# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001) -# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) -# config.network_fn = lambda: DisjointActorCriticNet( -# # lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=False, action_gate=F.tanh, action_scale=2.0), -# lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=True), -# lambda: GaussianCriticNet(task.state_dim)) -# config.policy_fn = lambda: GaussianPolicy() -# config.worker = ContinuousAdvantageActorCritic -# config.discount = 0.99 -# config.num_workers = 8 -# config.update_interval = 20 -# config.test_interval = 1 -# config.test_repetitions = 1 -# config.entropy_weight = 0 -# config.gradient_clip = 40 -# config.logger = Logger('./log', logger) -# agent = AsyncAgent(config) -# agent.run() -# -# def p3o_continuous(): -# config = Config() -# config.task_fn = lambda: Pendulum() -# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') -# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') -# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') -# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') -# # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1') -# task = config.task_fn() -# config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim, -# gpu=-1, unit_std=True) -# config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=-1) -# config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn) -# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) -# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) -# -# config.policy_fn = lambda: GaussianPolicy() -# config.replay_fn = lambda: GeneralReplay(memory_size=2048, batch_size=2048) -# config.worker = ProximalPolicyOptimization -# config.discount = 0.99 -# config.gae_tau = 0.97 -# config.num_workers = 6 -# config.test_interval = 1 -# config.test_repetitions = 1 -# config.entropy_weight = 0 -# config.gradient_clip = 20 -# config.rollout_length = 10000 -# config.optimize_epochs = 1 -# config.ppo_ratio_clip = 0.2 -# config.logger = Logger('./log', logger) -# agent = AsyncAgent(config) -# agent.run() -# -# def d3pg_continuous(): -# config = Config() -# config.task_fn = lambda: Pendulum() -# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') -# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') -# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') -# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') -# # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') -# task = config.task_fn() -# config.actor_network_fn = lambda: DeterministicActorNet( -# task.state_dim, task.action_dim, F.tanh, 2, non_linear=F.relu, batch_norm=False) -# config.critic_network_fn = lambda: DeterministicCriticNet( -# task.state_dim, task.action_dim, non_linear=F.relu, batch_norm=False) -# config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn) -# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4) -# config.critic_optimizer_fn =\ -# lambda params: torch.optim.Adam(params, lr=1e-4) -# config.replay_fn = lambda: SharedReplay(memory_size=1000000, batch_size=64, -# state_shape=(task.state_dim, ), action_shape=(task.action_dim, )) -# config.discount = 0.99 -# config.random_process_fn = \ -# lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2, -# n_steps_annealing=100000) -# config.worker = DeterministicPolicyGradient -# config.num_workers = 6 -# config.min_memory_size = 50 -# config.target_network_mix = 0.001 -# config.test_interval = 500 -# config.test_repetitions = 1 -# config.gradient_clip = 20 -# config.logger = Logger('./log', logger) -# agent = AsyncAgent(config) -# agent.run() +## continuous control + +def ppo_continuous(): + config = Config() + config.num_workers = 5 + task_fn = lambda log_dir: Pendulum(log_dir=log_dir) + # task_fn = lambda log_dir: Roboschool('RoboschoolInvertedPendulum-v1', log_dir=log_dir) + # task_fn = lambda log_dir: Roboschool('RoboschoolAnt-v1', log_dir=log_dir) + config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, log_dir=get_default_log_dir(ppo_continuous.__name__)) + config.actor_network_fn = lambda state_dim, action_dim: GaussianActorNet(state_dim, action_dim) + config.critic_network_fn = lambda state_dim, action_dim: GaussianCriticNet(state_dim) + config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) + config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) + config.discount = 0.99 + config.use_gae = True + config.gae_tau = 0.97 + config.gradient_clip = 0.5 + config.rollout_length = 20 + config.optimize_epochs = 4 + config.ppo_ratio_clip = 0.2 + config.logger = Logger('./log', logger) + run_iterations(PPOAgent(config)) def ddpg_continuous(): config = Config() config.task_fn = lambda: Pendulum() - # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2') - # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2') - # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2') # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1') # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1') # config.task_fn = lambda: Roboschool('RoboschoolHopper-v1') @@ -346,9 +273,8 @@ if __name__ == '__main__': # n_step_dqn_pixel_atari('BreakoutNoFrameskip-v4') # dqn_ram_atari('Breakout-ramNoFrameskip-v4') - ddpg_continuous() + # ddpg_continuous() + ppo_continuous() - # dqn_pixel_atari('BreakoutNoFrameskip-v4') - # dqn_ram_atari('Pong-ramNoFrameskip-v4') # acvp.train('PongNoFrameskip-v4') diff --git a/network/continuous_action_network.py b/network/continuous_action_network.py index cbdeefc..6e7f417 100644 --- a/network/continuous_action_network.py +++ b/network/continuous_action_network.py @@ -92,7 +92,7 @@ class GaussianActorNet(nn.Module, BasicNet): action_scale=1, action_gate=F.tanh, gpu=-1, - unit_std=True, + # unit_std=True, hidden_size=64, non_linear=F.tanh): super(GaussianActorNet, self).__init__() @@ -100,12 +100,8 @@ class GaussianActorNet(nn.Module, BasicNet): self.fc2 = nn.Linear(hidden_size, hidden_size) self.action_mean = nn.Linear(hidden_size, action_dim) - if unit_std: - self.action_log_std = nn.Parameter(torch.zeros(1, action_dim)) - else: - self.action_std = nn.Linear(hidden_size, action_dim) + self.action_log_std = nn.Parameter(torch.zeros(1, action_dim)) - self.unit_std = unit_std self.action_scale = action_scale self.action_gate = action_gate self.non_linear = non_linear @@ -119,24 +115,20 @@ class GaussianActorNet(nn.Module, BasicNet): mean = self.action_mean(phi) if self.action_gate is not None: mean = self.action_scale * self.action_gate(mean) - if self.unit_std: - log_std = self.action_log_std.expand_as(mean) - std = log_std.exp() - else: - std = F.softplus(self.action_std(phi)) + 1e-5 - log_std = std.log() + log_std = self.action_log_std.expand_as(mean) + std = log_std.exp() return mean, std, log_std def predict(self, x): return self.forward(x) - def log_density(self, x, mean, log_std, std): - var = std.pow(2) - log_density = -(x - mean).pow(2) / (2 * var + 1e-5) - 0.5 * torch.log(2 * Variable(torch.FloatTensor([np.pi])).expand_as(x)) - log_std - return log_density.sum(1) - - def entropy(self, std): - return 0.5 * (1 + (2 * std.pow(2) * np.pi + 1e-5).log()).sum(1).mean() + # def log_density(self, x, mean, log_std, std): + # var = std.pow(2) + # log_density = -(x - mean).pow(2) / (2 * var + 1e-5) - 0.5 * torch.log(2 * Variable(torch.FloatTensor([np.pi])).expand_as(x)) - log_std + # return log_density.sum(1) + # + # def entropy(self, std): + # return 0.5 * (1 + (2 * std.pow(2) * np.pi + 1e-5).log()).sum(1).mean() class GaussianCriticNet(nn.Module, BasicNet): def __init__(self, diff --git a/utils/config.py b/utils/config.py index a9e7bce..d3b0038 100644 --- a/utils/config.py +++ b/utils/config.py @@ -57,3 +57,4 @@ class Config: self.categorical_n_atoms = 51 self.num_quantiles = 10 self.gaussian_noise_scale = 0.3 + self.optimization_epochs = 4 diff --git a/utils/normalizer.py b/utils/normalizer.py index 9e690d6..99cb98f 100644 --- a/utils/normalizer.py +++ b/utils/normalizer.py @@ -13,6 +13,17 @@ class Normalizer: self.n = 1.0 def __call__(self, x): + if np.isscalar(x) or len(x.shape) == 1: + return self.nomalize_single(x) + elif len(x.shape) == 2: + new_x = np.zeros(x.shape) + for i in range(x.shape[0]): + new_x[i] = self.nomalize_single(x[i]) + return new_x + else: + assert 'Unsupported Shape' + + def nomalize_single(self, x): is_scalar = np.isscalar(x) if is_scalar: x = np.asarray([x])