mirror of
https://github.com/wassname/DeepRL.git
synced 2026-08-31 11:41:31 +08:00
Discrete PPO cartpole
This commit is contained in:
+3
-4
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user