Separate policy and value net

This commit is contained in:
Antonin RAFFIN 2019-09-21 18:12:06 +02:00
parent 2469ff3859
commit 70e1d673a9

View file

@ -27,6 +27,7 @@ class PPOPolicy(BasePolicy):
'activation_fn': self.activation_fn 'activation_fn': self.activation_fn
} }
self.shared_net = None self.shared_net = None
self.pi_net, self.vf_net = None, None
self._build(learning_rate) self._build(learning_rate)
@staticmethod @staticmethod
@ -36,16 +37,24 @@ class PPOPolicy(BasePolicy):
module.bias.data.fill_(0.0) module.bias.data.fill_(0.0)
def _build(self, learning_rate): def _build(self, learning_rate):
# TODO: support non-shared network # TODO: support shared network
shared_net = create_mlp(self.state_dim, output_dim=-1, net_arch=self.net_arch, activation_fn=self.activation_fn) # shared_net = create_mlp(self.state_dim, output_dim=-1, net_arch=self.net_arch, activation_fn=self.activation_fn)
self.shared_net = nn.Sequential(*shared_net).to(self.device) # self.shared_net = nn.Sequential(*shared_net).to(self.device)
pi_net = create_mlp(self.state_dim, output_dim=-1, net_arch=self.net_arch, activation_fn=self.activation_fn)
self.pi_net = nn.Sequential(*pi_net).to(self.device)
vf_net = create_mlp(self.state_dim, output_dim=-1, net_arch=self.net_arch, activation_fn=self.activation_fn)
self.vf_net = nn.Sequential(*vf_net).to(self.device)
self.actor_net = nn.Linear(self.net_arch[-1], self.action_dim) self.actor_net = nn.Linear(self.net_arch[-1], self.action_dim)
self.value_net = nn.Linear(self.net_arch[-1], 1) self.value_net = nn.Linear(self.net_arch[-1], 1)
self.log_std = nn.Parameter(th.zeros(self.action_dim)) self.log_std = nn.Parameter(th.zeros(self.action_dim))
# Init weights: use orthogonal initialization # Init weights: use orthogonal initialization
for module in [self.shared_net, self.actor_net, self.value_net]: for module in [self.pi_net, self.vf_net, self.actor_net, self.value_net]:
# Values from stable-baselines check why # Values from stable-baselines check why
gain = { gain = {
self.pi_net: np.sqrt(2),
self.vf_net: np.sqrt(2),
self.shared_net: np.sqrt(2), self.shared_net: np.sqrt(2),
self.actor_net: 0.01, self.actor_net: 0.01,
self.value_net: 1 self.value_net: 1
@ -56,12 +65,19 @@ class PPOPolicy(BasePolicy):
def forward(self, state, deterministic=False): def forward(self, state, deterministic=False):
state = th.FloatTensor(state).to(self.device) state = th.FloatTensor(state).to(self.device)
latent = self.shared_net(state) latent_pi, latent_vf = self._get_latent(state)
value = self.value_net(latent) value = self.value_net(latent_vf)
action, action_distribution = self._get_action_dist_from_latent(latent, deterministic=deterministic) action, action_distribution = self._get_action_dist_from_latent(latent_pi, deterministic=deterministic)
log_prob = self._get_log_prob(action_distribution, action) log_prob = self._get_log_prob(action_distribution, action)
return action, value, log_prob return action, value, log_prob
def _get_latent(self, state):
if self.shared_net is not None:
latent = self.shared_net(state)
return latent, latent
else:
return self.pi_net(state), self.vf_net(state)
def _get_action_dist_from_latent(self, latent, deterministic=False): def _get_action_dist_from_latent(self, latent, deterministic=False):
mean_actions = self.actor_net(latent) mean_actions = self.actor_net(latent)
action_std = th.ones_like(mean_actions) * self.log_std.exp() action_std = th.ones_like(mean_actions) * self.log_std.exp()
@ -83,15 +99,15 @@ class PPOPolicy(BasePolicy):
return log_prob return log_prob
def actor_forward(self, state, deterministic=False): def actor_forward(self, state, deterministic=False):
latent = self.shared_net(state) latent_pi, _ = self._get_latent(state)
action, _ = self._get_action_dist_from_latent(latent, deterministic=deterministic) action, _ = self._get_action_dist_from_latent(latent_pi, deterministic=deterministic)
return action.detach().cpu().numpy() return action.detach().cpu().numpy()
def get_policy_stats(self, state, action): def get_policy_stats(self, state, action):
latent = self.shared_net(state) latent_pi, latent_vf = self._get_latent(state)
_, action_distribution = self._get_action_dist_from_latent(latent) _, action_distribution = self._get_action_dist_from_latent(latent_pi)
log_prob = self._get_log_prob(action_distribution, action) log_prob = self._get_log_prob(action_distribution, action)
value = self.value_net(latent) value = self.value_net(latent_vf)
return value, log_prob, action_distribution.entropy() return value, log_prob, action_distribution.entropy()
def value_forward(self): def value_forward(self):