Update PPO

This commit is contained in:
Shangtong Zhang
2018-04-08 09:14:35 -06:00
parent 0f3f49252b
commit 45587b9cce
4 changed files with 83 additions and 41 deletions
+32 -19
View File
@@ -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
+26 -14
View File
@@ -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))
+24 -8
View File
@@ -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))
+1
View File
@@ -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