diff --git a/README.md b/README.md index 87f098a..3147aca 100644 --- a/README.md +++ b/README.md @@ -1,20 +1,30 @@ -# pytorch-a2c +# pytorch-a2c-ppo -This is a PyTorch implementation of Advantage Actor Critic (A2C), a synchronous deterministic version of A3C ["Asynchronous Methods for Deep Reinforcement Learning"](https://arxiv.org/pdf/1602.01783v1.pdf). Also see [the OpenAI post](https://blog.openai.com/baselines-acktr-a2c/) (section A2C and A3C) for more information. +This is a PyTorch implementation of Advantage Actor Critic (A2C), a synchronous deterministic version of A3C ["Asynchronous Methods for Deep Reinforcement Learning"](https://arxiv.org/pdf/1602.01783v1.pdf) and [Proximal Policy Optimization (PPO)](https://arxiv.org/pdf/1707.06347.pdf). Also see the OpenAI posts: [A2C/A3C](https://blog.openai.com/baselines-acktr-a2c/) and [PPO](https://blog.openai.com/openai-baselines-ppo/) for more information. -This implementation is inspired by the [OpenAI A2C baseline](https://github.com/openai/baselines/tree/master/baselines/a2c). It uses the same hyper parameters and the model since they were well tuned for Atari games. +This implementation is inspired by the OpenAI baselines for [A2C](https://github.com/openai/baselines/tree/master/baselines/a2c) and [PPO](https://github.com/openai/baselines/tree/master/baselines/ppo1). It uses the same hyper parameters and the model since they were well tuned for Atari games. ## Contributions Contributions are very welcome. If you know how to make this code better, don't hesitate to send a pull request. ## Usage + +### A2C ``` python main.py --env-name "PongNoFrameskip-v4" ``` +### PPO + +``` +python main.py --env-name "PongNoFrameskip-v4" --algo ppo --use-gae --num-processes 8 --num-steps 256 --vis-interval 1 --log-interval 1 +``` + ## Results +### A2C + ![BreakoutNoFrameskip-v4](imgs/breakout.png) ![SeaquestNoFrameskip-v4](imgs/seaquest.png) @@ -23,4 +33,6 @@ python main.py --env-name "PongNoFrameskip-v4" ![beamriderNoFrameskip-v4](imgs/beamrider.png) -More results coming soon. +### PPO + +Coming soon. diff --git a/main.py b/main.py index 9e49385..e0ef881 100755 --- a/main.py +++ b/main.py @@ -9,13 +9,16 @@ import torch.nn as nn import torch.nn.functional as F import torch.optim as optim from torch.autograd import Variable +from torch.utils.data.sampler import BatchSampler, SubsetRandomSampler from baselines.common.vec_env.subproc_vec_env import SubprocVecEnv from envs import make_env from model import ActorCritic from vizualize_atari import visdom_plot -parser = argparse.ArgumentParser(description='A2C') +parser = argparse.ArgumentParser(description='RL') +parser.add_argument('--algo', default='a2c', + help='algorithm to use: a2c | ppo') parser.add_argument('--lr', type=float, default=7e-4, help='learning rate (default: 7e-4)') parser.add_argument('--eps', type=float, default=1e-5, @@ -40,6 +43,12 @@ parser.add_argument('--num-processes', type=int, default=16, help='how many training CPU processes to use (default: 16)') parser.add_argument('--num-steps', type=int, default=5, help='number of forward steps in A2C (default: 5)') +parser.add_argument('--ppo-epoch', type=int, default=4, + help='number of ppo epochs (default: 4)') +parser.add_argument('--batch-size', type=int, default=64, + help='ppo batch size (default: 64)') +parser.add_argument('--clip-param', type=float, default=0.2, + help='ppo clip parameter (default: 0.2)') parser.add_argument('--num-stack', type=int, default=4, help='number of frames to stack (default: 4)') parser.add_argument('--log-interval', type=int, default=10, @@ -59,6 +68,10 @@ parser.add_argument('--no-vis', action='store_true', default=False, args = parser.parse_args() + +assert args.algo in ['a2c', 'ppo'] +if args.algo == 'ppo': + assert args.num_processes * args.num_steps % args.batch_size == 0 args.cuda = not args.no_cuda and torch.cuda.is_available() args.vis = not args.no_vis @@ -94,11 +107,16 @@ def main(): ]) actor_critic = ActorCritic(envs.observation_space.shape[0] * args.num_stack, envs.action_space) + if args.algo == 'ppo': + actor_critic = nn.DataParallel(actor_critic) if args.cuda: actor_critic.cuda() - optimizer = optim.RMSprop(actor_critic.parameters(), args.lr, eps=args.eps, alpha=args.alpha) + if args.algo == 'a2c': + optimizer = optim.RMSprop(actor_critic.parameters(), args.lr, eps=args.eps, alpha=args.alpha) + elif args.algo == 'ppo': + optimizer = optim.Adam(actor_critic.parameters(), eps=args.eps) obs_shape = envs.observation_space.shape obs_shape = (obs_shape[0] * args.num_stack, obs_shape[1], obs_shape[2]) @@ -117,6 +135,7 @@ def main(): rewards = torch.zeros(args.num_steps, args.num_processes, 1) value_preds = torch.zeros(args.num_steps + 1, args.num_processes, 1) + old_log_probs = torch.zeros(args.num_steps, args.num_processes, envs.action_space.n) returns = torch.zeros(args.num_steps + 1, args.num_processes, 1) actions = torch.LongTensor(args.num_steps, args.num_processes) @@ -131,6 +150,7 @@ def main(): current_state = current_state.cuda() rewards = rewards.cuda() value_preds = value_preds.cuda() + old_log_probs = old_log_probs.cuda() returns = returns.cuda() actions = actions.cuda() masks = masks.cuda() @@ -163,6 +183,7 @@ def main(): update_current_state(state) states[step + 1].copy_(current_state) value_preds[step].copy_(value.data) + old_log_probs[step].copy_(log_probs) rewards[step].copy_(reward) masks[step].copy_(torch.from_numpy(np_masks).unsqueeze(1)) @@ -177,7 +198,6 @@ def main(): for step in reversed(range(args.num_steps)): delta = rewards[step] + args.gamma * value_preds[step + 1] * masks[step] - value_preds[step] gae = delta + args.gamma * args.tau * masks[step] * gae - returns[step] = gae + value_preds[step] else: returns[-1] = actor_critic(Variable(states[-1], volatile=True))[0].data @@ -185,35 +205,68 @@ def main(): returns[step] = returns[step + 1] * \ args.gamma * masks[step] + rewards[step] + if args.algo == 'a2c': + # Reshape to do in a single forward pass for all steps + values, logits = actor_critic(Variable(states[:-1].view(-1, *states.size()[-3:]))) + log_probs = F.log_softmax(logits) - # Reshape to do in a single forward pass for all steps - values, logits = actor_critic(Variable(states[:-1].view(-1, *states.size()[-3:]))) - log_probs = F.log_softmax(logits) - probs = F.softmax(logits) + # Unreshape + logits_size = (args.num_steps, args.num_processes, logits.size(-1)) - # Unreshape - logits_size = (args.num_steps, args.num_processes, logits.size(-1)) + log_probs = F.log_softmax(logits).view(logits_size) + probs = F.softmax(logits).view(logits_size) - log_probs = F.log_softmax(logits).view(logits_size) - probs = F.softmax(logits).view(logits_size) + values = values.view(args.num_steps, args.num_processes, 1) + logits = logits.view(logits_size) - values = values.view(args.num_steps, args.num_processes, 1) - logits = logits.view(logits_size) + action_log_probs = log_probs.gather(2, Variable(actions.unsqueeze(2))) - action_log_probs = log_probs.gather(2, Variable(actions.unsqueeze(2))) + dist_entropy = -(log_probs * probs).sum(-1).mean() - dist_entropy = -(log_probs * probs).sum(-1).mean() + advantages = Variable(returns[:-1]) - values + value_loss = advantages.pow(2).mean() - advantages = Variable(returns[:-1]) - values - value_loss = advantages.pow(2).mean() + action_loss = -(Variable(advantages.data) * action_log_probs).mean() - action_loss = -(Variable(advantages.data) * action_log_probs).mean() + optimizer.zero_grad() + (value_loss * args.value_loss_coef + action_loss - dist_entropy * args.entropy_coef).backward() - optimizer.zero_grad() - (value_loss * args.value_loss_coef + action_loss - dist_entropy * args.entropy_coef).backward() + nn.utils.clip_grad_norm(actor_critic.parameters(), args.max_grad_norm) + optimizer.step() + elif args.algo == 'ppo': + advantages = returns[:-1] - value_preds[:-1] + advantages = (advantages - advantages.mean()) / advantages.std() + for _ in range(args.ppo_epoch): + sampler = BatchSampler(SubsetRandomSampler(range(args.num_processes * args.num_steps)), args.batch_size * args.num_processes, drop_last=False) + for indices in sampler: + states_batch = states[:-1].view(-1, *states.size()[-3:])[indices] + actions_batch = actions.view(-1, 1)[indices] + return_batch = returns[:-1].view(-1, 1)[indices] - nn.utils.clip_grad_norm(actor_critic.parameters(), args.max_grad_norm) - optimizer.step() + # Reshape to do in a single forward pass for all steps + values, logits = actor_critic(Variable(states_batch)) + log_probs = F.log_softmax(logits) + action_log_probs = log_probs.gather(1, Variable(actions_batch)) + + old_log_probs_batch = old_log_probs.view(-1, old_log_probs.size(-1))[indices] + old_action_log_probs = old_log_probs_batch.gather(1, actions_batch) + + ratio = torch.exp(action_log_probs - Variable(old_action_log_probs)) + adv_targ = Variable(advantages.view(-1, 1)[indices]) + surr1 = ratio * adv_targ + surr2 = ratio.clamp(1.0 - args.clip_param, 1.0 + args.clip_param) * adv_targ + action_loss = -torch.min(surr1, surr2).mean() # PPO's pessimistic surrogate (L^CLIP) + + log_probs = F.log_softmax(logits) + probs = F.softmax(logits) + + dist_entropy = -(log_probs * probs).sum(-1).mean() + + value_loss = (Variable(return_batch) - values).pow(2).mean() + + optimizer.zero_grad() + (value_loss + action_loss - dist_entropy * args.entropy_coef).backward() + optimizer.step() states[0].copy_(states[-1]) @@ -227,7 +280,7 @@ def main(): value_loss.data[0], action_loss.data[0])) if j % args.vis_interval == 0: - win = visdom_plot(viz, win, args.log_dir, args.env_name, 'a2c') + win = visdom_plot(viz, win, args.log_dir, args.env_name, args.algo) if __name__ == "__main__":