diff --git a/agent/PPO_agent.py b/agent/PPO_agent.py index bc87326..41adf1e 100644 --- a/agent/PPO_agent.py +++ b/agent/PPO_agent.py @@ -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() diff --git a/main.py b/main.py index 354baa3..7a35a48 100644 --- a/main.py +++ b/main.py @@ -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() diff --git a/network/base_network.py b/network/base_network.py index 9425a06..60ba6dd 100644 --- a/network/base_network.py +++ b/network/base_network.py @@ -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)