diff --git a/agent/PPO_agent.py b/agent/PPO_agent.py index 550c44e..8564360 100644 --- a/agent/PPO_agent.py +++ b/agent/PPO_agent.py @@ -30,11 +30,6 @@ class PPOWorker: # self.shared_state_shifter() - def normal_log_density(self, x, mean, log_std, std): - var = std.pow(2) - log_density = -(x - mean).pow(2) / (2 * var) - 0.5 * torch.log(2 * Variable(torch.FloatTensor([np.pi])).expand_as(x)) - log_std - return log_density.sum(1) - def rollout(self, deterministic=False): config = self.config replay = config.replay_fn() @@ -114,13 +109,15 @@ class PPOWorker: advantages = (advantages - advantages.mean().expand_as(advantages)) / advantages.std().expand_as(advantages) mean_old, std_old, log_std_old = actor_net_old.predict(states) - probs_old = self.normal_log_density(actions, mean_old, log_std_old, std_old) + probs_old = self.actor_net.log_density(actions, mean_old, log_std_old, std_old) mean, std, log_std = self.actor_net.predict(states) - probs = self.normal_log_density(actions, mean, log_std, std) + probs = self.actor_net.log_density(actions, mean, log_std, std) ratio = (probs - 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) + if config.entropy_weight: + policy_loss += config.entropy_weight * self.actor_net.kl_loss(std) v = self.critic_net.predict(states) value_loss = 0.5 * (returns - v).pow(2).mean() @@ -130,8 +127,6 @@ class PPOWorker: nn.utils.clip_grad_norm(self.critic_net.parameters(), config.gradient_clip) self.critic_opt.step() - # kl_loss = 0.5 * (1 + (2 * std.pow(2) * np.pi + 1e-5).log()).sum(1).mean() - # policy_loss += kl_loss actor_net_old.load_state_dict(self.actor_net.state_dict()) self.actor_opt.zero_grad() diff --git a/main.py b/main.py index a442af2..86a012d 100644 --- a/main.py +++ b/main.py @@ -308,8 +308,8 @@ def ppo_pendulum(): # config.task_fn = lambda: BipedalWalker() # config.reward_shift_fn = lambda reward: reward / 10 task = config.task_fn() - config.actor_network_fn = lambda: PPOActorNet(task.state_dim, task.action_dim) - config.critic_network_fn = lambda: PPOCriticNet(task.state_dim) + config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim) + config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim) config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) @@ -320,10 +320,9 @@ def ppo_pendulum(): config.gae_tau = 0.97 config.max_episode_length = 200 config.num_workers = None - config.update_interval = None config.test_interval = None config.test_repetitions = None - config.entropy_weight = 0.001 + config.entropy_weight = 0 config.gradient_clip = 40 config.rollout_length = 10000 config.optimize_epochs = 10 @@ -339,11 +338,11 @@ if __name__ == '__main__': # dqn_cart_pole() # async_cart_pole() # a3c_cart_pole() - # a3c_pendulum() + a3c_pendulum() # a3c_walker() # ddpg_pendulum() # ddpg_walker() - ppo_pendulum() + # ppo_pendulum() # dqn_fruit() # hrdqn_fruit() diff --git a/network/continuous_action_network.py b/network/continuous_action_network.py index 49cb389..9309436 100644 --- a/network/continuous_action_network.py +++ b/network/continuous_action_network.py @@ -142,15 +142,20 @@ class DDPGCriticNet(nn.Module, BasicNet): def predict(self, x, action): return self.forward(x, action) -class PPOActorNet(nn.Module, BasicNet): - def __init__(self, state_dim, action_dim, action_scale=1.0, action_gate=None, gpu=False): - super(PPOActorNet, self).__init__() +class GaussianActorNet(nn.Module, BasicNet): + def __init__(self, state_dim, action_dim, action_scale=1.0, action_gate=None, gpu=False, unit_std=True): + super(GaussianActorNet, self).__init__() hidden_size = 64 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.action_log_std = nn.Parameter(torch.zeros(1, action_dim)) + if unit_std: + self.action_log_std = nn.Parameter(torch.zeros(1, action_dim)) + else: + self.action_std = nn.Linear(hidden_size, action_dim) + + self.unit_std = unit_std self.action_scale = action_scale self.action_gate = action_gate @@ -163,17 +168,28 @@ class PPOActorNet(nn.Module, BasicNet): mean = self.action_mean(phi) if self.action_gate is not None: mean = self.action_scale * self.action_gate(mean) - log_std = self.action_log_std.expand_as(mean) - std = log_std.exp() + if self.unit_std: + log_std = self.action_log_std.expand_as(mean) + std = log_std.exp() + else: + std = F.softplus(self.fc_std(x) + 1e-5) + log_std = std.log() return mean, std, log_std def predict(self, x): return self.forward(x) + def log_density(self, x, mean, log_std, std): + var = std.pow(2) + log_density = -(x - mean).pow(2) / (2 * var) - 0.5 * torch.log(2 * Variable(torch.FloatTensor([np.pi])).expand_as(x)) - log_std + return log_density.sum(1) -class PPOCriticNet(nn.Module, BasicNet): + def kl_loss(self, std): + return 0.5 * (1 + (2 * std.pow(2) * np.pi + 1e-5).log()).sum(1).mean() + +class GaussianCriticNet(nn.Module, BasicNet): def __init__(self, state_dim, gpu=False): - super(PPOCriticNet, self).__init__() + super(GaussianCriticNet, self).__init__() hidden_size = 64 self.fc1 = nn.Linear(state_dim, hidden_size) self.fc2 = nn.Linear(hidden_size, hidden_size)