mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-15 22:10:25 +00:00
Separate policy and value net
This commit is contained in:
parent
2469ff3859
commit
70e1d673a9
1 changed files with 28 additions and 12 deletions
|
|
@ -27,6 +27,7 @@ class PPOPolicy(BasePolicy):
|
|||
'activation_fn': self.activation_fn
|
||||
}
|
||||
self.shared_net = None
|
||||
self.pi_net, self.vf_net = None, None
|
||||
self._build(learning_rate)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -36,16 +37,24 @@ class PPOPolicy(BasePolicy):
|
|||
module.bias.data.fill_(0.0)
|
||||
|
||||
def _build(self, learning_rate):
|
||||
# TODO: support non-shared network
|
||||
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)
|
||||
# TODO: support shared network
|
||||
# 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)
|
||||
|
||||
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.value_net = nn.Linear(self.net_arch[-1], 1)
|
||||
self.log_std = nn.Parameter(th.zeros(self.action_dim))
|
||||
# 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
|
||||
gain = {
|
||||
self.pi_net: np.sqrt(2),
|
||||
self.vf_net: np.sqrt(2),
|
||||
self.shared_net: np.sqrt(2),
|
||||
self.actor_net: 0.01,
|
||||
self.value_net: 1
|
||||
|
|
@ -56,12 +65,19 @@ class PPOPolicy(BasePolicy):
|
|||
|
||||
def forward(self, state, deterministic=False):
|
||||
state = th.FloatTensor(state).to(self.device)
|
||||
latent = self.shared_net(state)
|
||||
value = self.value_net(latent)
|
||||
action, action_distribution = self._get_action_dist_from_latent(latent, deterministic=deterministic)
|
||||
latent_pi, latent_vf = self._get_latent(state)
|
||||
value = self.value_net(latent_vf)
|
||||
action, action_distribution = self._get_action_dist_from_latent(latent_pi, deterministic=deterministic)
|
||||
log_prob = self._get_log_prob(action_distribution, action)
|
||||
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):
|
||||
mean_actions = self.actor_net(latent)
|
||||
action_std = th.ones_like(mean_actions) * self.log_std.exp()
|
||||
|
|
@ -83,15 +99,15 @@ class PPOPolicy(BasePolicy):
|
|||
return log_prob
|
||||
|
||||
def actor_forward(self, state, deterministic=False):
|
||||
latent = self.shared_net(state)
|
||||
action, _ = self._get_action_dist_from_latent(latent, deterministic=deterministic)
|
||||
latent_pi, _ = self._get_latent(state)
|
||||
action, _ = self._get_action_dist_from_latent(latent_pi, deterministic=deterministic)
|
||||
return action.detach().cpu().numpy()
|
||||
|
||||
def get_policy_stats(self, state, action):
|
||||
latent = self.shared_net(state)
|
||||
_, action_distribution = self._get_action_dist_from_latent(latent)
|
||||
latent_pi, latent_vf = self._get_latent(state)
|
||||
_, action_distribution = self._get_action_dist_from_latent(latent_pi)
|
||||
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()
|
||||
|
||||
def value_forward(self):
|
||||
|
|
|
|||
Loading…
Reference in a new issue