Discrete PPO cartpole

This commit is contained in:
Shangtong Zhang
2018-04-10 23:37:12 -06:00
parent 015c7753ea
commit bb5ef4fbfc
3 changed files with 62 additions and 5 deletions
+3 -4
View File
@@ -82,18 +82,17 @@ class PPOAgent(BaseAgent):
sampled_returns = returns[batch_indices]
sampled_advantages = advantages[batch_indices]
_, log_probs, _, values = self.network.predict(sampled_states, sampled_actions)
_, log_probs, entropy_loss, values = self.network.predict(sampled_states, sampled_actions)
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)
policy_loss = -torch.min(obj, obj_clipped).mean(0) + config.entropy_weight * entropy_loss.mean(0)
value_loss = 0.5 * (sampled_returns - values).pow(2).mean()
self.network.zero_grad()
policy_loss.backward()
value_loss.backward()
(policy_loss + value_loss).backward()
nn.utils.clip_grad_norm(self.network.parameters(), config.gradient_clip)
self.network.step()
+24 -1
View File
@@ -94,6 +94,28 @@ def n_step_dqn_cart_pole():
config.logger = Logger('./log', logger)
run_iterations(NStepDQNAgent(config))
def ppo_cart_pole():
config = Config()
task_fn = lambda log_dir: ClassicalControl('CartPole-v0', max_steps=200, log_dir=log_dir)
config.num_workers = 5
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001)
network_fn = lambda state_dim, action_dim: ActorCriticFCNet(state_dim, 64, action_dim)
config.network_fn = lambda state_dim, action_dim: \
DiscreteActorCriticWrapper(state_dim, action_dim, network_fn, optimizer_fn)
config.discount = 0.99
config.logger = Logger('./log', logger)
config.use_gae = True
config.gae_tau = 0.95
config.entropy_weight = 0.01
config.gradient_clip = 0.5
config.rollout_length = 128
config.optimization_epochs = 10
config.num_mini_batches = 4
config.ppo_ratio_clip = 0.2
config.iteration_log_interval = 1
run_iterations(PPOAgent(config))
## Atari games
def dqn_pixel_atari(name):
@@ -327,6 +349,7 @@ if __name__ == '__main__':
# categorical_dqn_cart_pole()
# quantile_regression_dqn_cart_pole()
# n_step_dqn_cart_pole()
ppo_cart_pole()
# dqn_pixel_atari('BreakoutNoFrameskip-v4')
# a2c_pixel_atari('BreakoutNoFrameskip-v4')
@@ -336,7 +359,7 @@ if __name__ == '__main__':
# dqn_ram_atari('Breakout-ramNoFrameskip-v4')
# ddpg_continuous()
ppo_continuous()
# ppo_continuous()
# action_conditional_video_prediction()
+35
View File
@@ -186,3 +186,38 @@ class ContinuousActorCriticWrapper:
def load_state_dict(self, state_dicts):
self.actor.load_state_dict(state_dicts[0])
self.critic.load_state_dict(state_dicts[1])
class DiscreteActorCriticWrapper:
def __init__(self, state_dim, action_dim, network_fn, opt_fn):
self.network = network_fn(state_dim, action_dim)
self.opt = opt_fn(self.network.parameters())
def predict(self, state, action=None):
prob, log_prob, value = self.network.predict(state)
entropy_loss = torch.sum(prob * log_prob, dim=1, keepdim=True)
dist = torch.distributions.Categorical(prob)
if action is None:
action = dist.sample()
log_prob = dist.log_prob(action).unsqueeze(1)
return action, log_prob, entropy_loss, value
def variable(self, x, dtype=torch.FloatTensor):
return self.network.variable(x, dtype)
def tensor(self, x, dtype=torch.FloatTensor):
return self.network.tensor(x, dtype)
def zero_grad(self):
self.opt.zero_grad()
def parameters(self):
return self.network.parameters()
def step(self):
self.opt.step()
def state_dict(self):
return self.network.state_dict()
def load_state_dict(self, state_dicts):
self.network.load_state_dict(state_dicts)