diff --git a/agent/PPO_agent.py b/agent/PPO_agent.py index 28b5684..98b724f 100644 --- a/agent/PPO_agent.py +++ b/agent/PPO_agent.py @@ -78,28 +78,41 @@ class PPOAgent(BaseAgent): 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) + returns = Variable(returns) - 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) + batcher = Batcher(states.size(0) // config.num_mini_batches, [np.arange(states.size(0))]) + for _ in range(config.optimization_epochs): + batcher.shuffle() + while not batcher.end(): + batch_indices = batcher.next_batch()[0] + batch_indices = self.actor.variable(batch_indices, torch.LongTensor) + sampled_states = states[batch_indices] + sampled_actions = actions[batch_indices] + sampled_log_probs_old = log_probs_old[batch_indices] + sampled_returns = returns[batch_indices] + sampled_advantages = advantages[batch_indices] - v = self.critic.predict(states) - value_loss = 0.5 * (Variable(returns) - v).pow(2).mean() + mean, std, log_std = self.actor.predict(sampled_states) + dist = torch.distributions.Normal(mean, std) + log_probs = dist.log_prob(sampled_actions) + log_probs = torch.sum(log_probs, dim=1, keepdim=True) + ratio = (log_probs - sampled_log_probs_old).exp() + obj = ratio * sampled_advantages + obj_clipped = ratio.clamp(1.0 - self.config.ppo_ratio_clip, + 1.0 + self.config.ppo_ratio_clip) * sampled_advantages + policy_loss = -torch.min(obj, obj_clipped).mean(0) - 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() + v = self.critic.predict(sampled_states) + value_loss = 0.5 * (sampled_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/main.py b/main.py index 8e925c3..26802e5 100644 --- a/main.py +++ b/main.py @@ -186,14 +186,15 @@ def n_step_dqn_pixel_atari(name): config.num_workers = 16 config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, log_dir=get_default_log_dir(n_step_dqn_pixel_atari.__name__)) - config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01) + config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=1e-4, alpha=0.99, eps=1e-5) config.network_fn = lambda state_dim, action_dim: ConvNet(config.history_length, action_dim, gpu=3) - config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1) + config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.05) config.state_normalizer = ImageNormalizer() config.reward_normalizer = SignNormalizer() config.discount = 0.99 config.target_network_update_freq = 10000 config.rollout_length = 5 + config.gradient_clip = 5 config.logger = Logger('./log', logger) run_iterations(NStepDQNAgent(config)) @@ -220,25 +221,29 @@ def dqn_ram_atari(name): def ppo_continuous(): config = Config() - config.num_workers = 16 + config.num_workers = 1 # 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) # task_fn = lambda log_dir: Roboschool('RoboschoolReacher-v1', log_dir=log_dir) task_fn = lambda log_dir: Roboschool('RoboschoolHopper-v1', log_dir=log_dir) + # task_fn = lambda log_dir: DMControl('cartpole', 'balance', log_dir=log_dir) + # task_fn = lambda log_dir: DMControl('hopper', 'hop', 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, 1e-4) - config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 1e-4) - config.state_normalizer = RunningStatsNormalizer() + config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 3e-4, eps=1e-5) + config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 3e-4, eps=1e-5) + # config.state_normalizer = RunningStatsNormalizer() config.discount = 0.99 config.use_gae = True - config.gae_tau = 0.97 + config.gae_tau = 0.95 config.gradient_clip = 0.5 - config.rollout_length = 20 - config.optimize_epochs = 5 + config.rollout_length = 2048 + config.optimization_epochs = 10 + config.num_mini_batches = 32 config.ppo_ratio_clip = 0.2 + config.iteration_log_interval = 1 config.logger = Logger('./log', logger) run_iterations(PPOAgent(config)) @@ -247,17 +252,19 @@ def ddpg_continuous(): log_dir = get_default_log_dir(ddpg_continuous.__name__) # config.task_fn = lambda: Pendulum(log_dir=log_dir) # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1', log_dir=log_dir) - config.task_fn = lambda: Roboschool('RoboschoolReacher-v1', log_dir=log_dir) - # config.task_fn = lambda: Roboschool('RoboschoolHopper-v1', log_dir=log_dir) + # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1', log_dir=log_dir) + config.task_fn = lambda: Roboschool('RoboschoolHopper-v1', log_dir=log_dir) # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1', log_dir=log_dir) # config.task_fn = lambda: Roboschool('RoboschoolWalker2d-v1', log_dir=log_dir) - config.task_fn = lambda: DMControl('cartpole', 'balance', log_dir=log_dir) + # config.task_fn = lambda: DMControl('cartpole', 'balance', log_dir=log_dir) + # config.task_fn = lambda: DMControl('finger', 'spin', log_dir=log_dir) config.actor_network_fn = lambda state_dim, action_dim: DeterministicActorNet(state_dim, action_dim) config.critic_network_fn = lambda state_dim, action_dim: DeterministicCriticNet(state_dim, action_dim) 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: HighDimActionReplay(memory_size=1000000, batch_size=64) config.discount = 0.99 + config.state_normalizer = RunningStatsNormalizer() config.random_process_fn = \ lambda action_dim: OrnsteinUhlenbeckProcess(size=action_dim, theta=0.15, sigma=0.3, n_steps_annealing=1000000) @@ -270,12 +277,17 @@ def ddpg_continuous(): def plot(): import matplotlib.pyplot as plt plotter = Plotter() + # name = 'log/ppo_continuous-180408-002056' + # plotter.plot_results([name]) + # plt.show() names = ['a2c_pixel_atari-180407-92711', 'categorical_dqn_pixel_atari-180407-094006', 'dqn_pixel_atari-180407-01414', 'quantile_regression_dqn_pixel_atari-180407-01604', - 'n_step_dqn_pixel_atari-180407-163421', - 'ppo_continuous-180407-111715'] + 'n_step_dqn_pixel_atari-180408-001104', + 'ppo_continuous-180408-002056', + 'ddpg_continuous-180407-234141' + ] for name in names: plotter.plot_results(['to_plot/%s' % (name)]) plt.savefig('images/%s.png' % (name)) diff --git a/network/continuous_action_network.py b/network/continuous_action_network.py index c1f6277..365c6d9 100644 --- a/network/continuous_action_network.py +++ b/network/continuous_action_network.py @@ -87,31 +87,36 @@ class GaussianActorNet(nn.Module, BasicNet): def __init__(self, state_dim, action_dim, - action_scale=1, - action_gate=F.tanh, gpu=-1, hidden_size=64, non_linear=F.tanh): super(GaussianActorNet, self).__init__() self.fc1 = nn.Linear(state_dim, hidden_size) self.fc2 = nn.Linear(hidden_size, hidden_size) - self.action_mean = nn.Linear(hidden_size, action_dim) + self.fc_action = nn.Linear(hidden_size, action_dim) self.action_log_std = nn.Parameter(torch.zeros(1, action_dim)) - self.action_scale = action_scale - self.action_gate = action_gate self.non_linear = non_linear + self.init_weights() BasicNet.__init__(self, gpu) + def init_weights(self): + bound = 3e-3 + nn.init.uniform(self.fc_action.weight.data, -bound, bound) + nn.init.constant(self.fc_action.bias.data, 0) + + nn.init.orthogonal(self.fc1.weight.data) + nn.init.constant(self.fc1.bias.data, 0) + nn.init.orthogonal(self.fc2.weight.data) + nn.init.constant(self.fc2.bias.data, 0) + def forward(self, x): x = self.variable(x) phi = self.non_linear(self.fc1(x)) phi = self.non_linear(self.fc2(phi)) - mean = self.action_mean(phi) - if self.action_gate is not None: - mean = self.action_scale * self.action_gate(mean) + mean = F.tanh(self.fc_action(phi)) log_std = self.action_log_std.expand_as(mean) std = log_std.exp() return mean, std, log_std @@ -130,8 +135,19 @@ class GaussianCriticNet(nn.Module, BasicNet): self.fc2 = nn.Linear(hidden_size, hidden_size) self.fc_value = nn.Linear(hidden_size, 1) self.non_linear = non_linear + self.init_weights() BasicNet.__init__(self, gpu) + def init_weights(self): + bound = 3e-3 + nn.init.uniform(self.fc_value.weight.data, -bound, bound) + nn.init.constant(self.fc_value.bias.data, 0) + + nn.init.orthogonal(self.fc1.weight.data) + nn.init.constant(self.fc1.bias.data, 0) + nn.init.orthogonal(self.fc2.weight.data) + nn.init.constant(self.fc2.bias.data, 0) + def forward(self, x): x = self.variable(x) phi = self.non_linear(self.fc1(x)) diff --git a/utils/config.py b/utils/config.py index 5c23b5c..38e0d7c 100644 --- a/utils/config.py +++ b/utils/config.py @@ -55,3 +55,4 @@ class Config: self.num_quantiles = 10 self.gaussian_noise_scale = 0.3 self.optimization_epochs = 4 + self.num_mini_batches = 32