Gaussian actor

This commit is contained in:
Shangtong Zhang
2017-10-04 10:15:28 -06:00
parent 27f54ed420
commit 3adceb5284
3 changed files with 33 additions and 23 deletions
+4 -9
View File
@@ -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()
+5 -6
View File
@@ -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()
+24 -8
View File
@@ -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)