mirror of
https://github.com/wassname/pytorch-a2c-ppo-acktr.git
synced 2026-09-12 12:41:11 +08:00
Add PPO
This commit is contained in:
@@ -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
|
||||
|
||||

|
||||
|
||||

|
||||
@@ -23,4 +33,6 @@ python main.py --env-name "PongNoFrameskip-v4"
|
||||
|
||||

|
||||
|
||||
More results coming soon.
|
||||
### PPO
|
||||
|
||||
Coming soon.
|
||||
|
||||
@@ -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__":
|
||||
|
||||
Reference in New Issue
Block a user