stable-baselines3/torchy_baselines/ppo/policies.py

121 lines
5.5 KiB
Python
Raw Normal View History

2019-09-21 14:48:51 +00:00
from functools import partial
2019-09-18 11:10:27 +00:00
import torch as th
import torch.nn as nn
2019-09-21 14:48:51 +00:00
import numpy as np
2019-09-18 11:10:27 +00:00
2019-09-18 13:35:17 +00:00
from torchy_baselines.common.policies import BasePolicy, register_policy, create_mlp
from torchy_baselines.common.distributions import make_proba_distribution, DiagGaussianDistribution, CategoricalDistribution
2019-09-18 11:10:27 +00:00
2019-09-24 12:53:03 +00:00
2019-09-18 11:10:27 +00:00
class PPOPolicy(BasePolicy):
def __init__(self, observation_space, action_space,
learning_rate=1e-3, net_arch=None, device='cpu',
2019-09-26 14:29:47 +00:00
activation_fn=nn.Tanh, adam_epsilon=1e-5, ortho_init=True):
2019-09-18 11:10:27 +00:00
super(PPOPolicy, self).__init__(observation_space, action_space, device)
2019-09-24 12:53:03 +00:00
self.obs_dim = self.observation_space.shape[0]
2019-09-18 11:10:27 +00:00
if net_arch is None:
net_arch = [64, 64]
self.net_arch = net_arch
self.activation_fn = activation_fn
2019-09-19 15:18:41 +00:00
self.adam_epsilon = adam_epsilon
2019-09-26 14:29:47 +00:00
self.ortho_init = ortho_init
2019-09-18 11:10:27 +00:00
self.net_args = {
2019-09-24 12:53:03 +00:00
'input_dim': self.obs_dim,
2019-09-18 11:10:27 +00:00
'output_dim': -1,
'net_arch': self.net_arch,
'activation_fn': self.activation_fn
}
self.shared_net = None
2019-09-21 16:12:06 +00:00
self.pi_net, self.vf_net = None, None
2019-09-22 10:52:49 +00:00
# Action distribution
self.action_dist = make_proba_distribution(action_space)
2019-09-18 11:10:27 +00:00
self._build(learning_rate)
def _build(self, learning_rate):
2019-09-21 16:12:06 +00:00
# TODO: support shared network
2019-09-24 12:53:03 +00:00
# shared_net = create_mlp(self.obs_dim, output_dim=-1, net_arch=self.net_arch, activation_fn=self.activation_fn)
2019-09-21 16:12:06 +00:00
# self.shared_net = nn.Sequential(*shared_net).to(self.device)
2019-09-24 12:53:03 +00:00
pi_net = create_mlp(self.obs_dim, output_dim=-1, net_arch=self.net_arch, activation_fn=self.activation_fn)
2019-09-21 16:12:06 +00:00
self.pi_net = nn.Sequential(*pi_net).to(self.device)
2019-09-24 12:53:03 +00:00
vf_net = create_mlp(self.obs_dim, output_dim=-1, net_arch=self.net_arch, activation_fn=self.activation_fn)
2019-09-21 16:12:06 +00:00
self.vf_net = nn.Sequential(*vf_net).to(self.device)
2019-09-22 10:52:49 +00:00
# self.action_net = nn.Linear(self.net_arch[-1], self.action_dim)
# self.log_std = nn.Parameter(th.zeros(self.action_dim))
if isinstance(self.action_dist, DiagGaussianDistribution):
self.action_net, self.log_std = self.action_dist.proba_distribution_net(latent_dim=self.net_arch[-1])
elif isinstance(self.action_dist, CategoricalDistribution):
self.action_net = self.action_dist.proba_distribution_net(latent_dim=self.net_arch[-1])
2019-09-18 11:10:27 +00:00
self.value_net = nn.Linear(self.net_arch[-1], 1)
2019-09-21 14:48:51 +00:00
# Init weights: use orthogonal initialization
2019-09-26 14:29:47 +00:00
# with small initial weight for the output
if self.ortho_init:
for module in [self.pi_net, self.vf_net, self.action_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.action_net: 0.01,
self.value_net: 1
}[module]
module.apply(partial(self.init_weights, gain=gain))
2019-09-21 14:48:51 +00:00
# TODO: support linear decay of the learning rate
2019-09-19 15:18:41 +00:00
self.optimizer = th.optim.Adam(self.parameters(), lr=learning_rate, eps=self.adam_epsilon)
2019-09-18 11:10:27 +00:00
# def get_action_dist_params(self, obs):
# latent_pi, _ = self._get_latent(obs)
# mean_actions = self.pi_net(latent_pi)
# if isinstance(self.action_dist, DiagGaussianDistribution):
# return {'mean_actions': mean_actions, 'log_std': self.log_std}
# elif isinstance(self.action_dist, CategoricalDistribution):
# return {'action_logits': mean_actions}
2019-09-24 12:53:03 +00:00
def forward(self, obs, deterministic=False):
if not isinstance(obs, th.Tensor):
obs = th.FloatTensor(obs).to(self.device)
latent_pi, latent_vf = self._get_latent(obs)
2019-09-21 16:12:06 +00:00
value = self.value_net(latent_vf)
action, action_distribution = self._get_action_dist_from_latent(latent_pi, deterministic=deterministic)
2019-09-22 10:52:49 +00:00
log_prob = action_distribution.log_prob(action)
# mean_actions, log_std = self.get_action_dist_params(obs)
# action, log_prob = self.action_dist.log_prob_from_params(**self.get_action_dist_params(obs))
2019-09-18 11:10:27 +00:00
return action, value, log_prob
2019-09-24 12:53:03 +00:00
def _get_latent(self, obs):
2019-09-21 16:12:06 +00:00
if self.shared_net is not None:
2019-09-24 12:53:03 +00:00
latent = self.shared_net(obs)
2019-09-21 16:12:06 +00:00
return latent, latent
else:
2019-09-24 12:53:03 +00:00
return self.pi_net(obs), self.vf_net(obs)
2019-09-21 16:12:06 +00:00
2019-09-19 09:43:15 +00:00
def _get_action_dist_from_latent(self, latent, deterministic=False):
2019-09-22 10:52:49 +00:00
mean_actions = self.action_net(latent)
if isinstance(self.action_dist, DiagGaussianDistribution):
return self.action_dist.proba_distribution(mean_actions, self.log_std, deterministic=deterministic)
elif isinstance(self.action_dist, CategoricalDistribution):
return self.action_dist.proba_distribution(mean_actions, deterministic=deterministic)
2019-09-19 09:43:15 +00:00
2019-09-24 12:53:03 +00:00
def actor_forward(self, obs, deterministic=False):
latent_pi, _ = self._get_latent(obs)
2019-09-21 16:12:06 +00:00
action, _ = self._get_action_dist_from_latent(latent_pi, deterministic=deterministic)
2019-09-19 09:43:15 +00:00
return action.detach().cpu().numpy()
2019-09-24 12:53:03 +00:00
def get_policy_stats(self, obs, action):
latent_pi, latent_vf = self._get_latent(obs)
2019-09-21 16:12:06 +00:00
_, action_distribution = self._get_action_dist_from_latent(latent_pi)
2019-09-22 10:52:49 +00:00
log_prob = action_distribution.log_prob(action)
2019-09-21 16:12:06 +00:00
value = self.value_net(latent_vf)
2019-09-19 09:43:15 +00:00
return value, log_prob, action_distribution.entropy()
2019-09-18 21:48:47 +00:00
2019-09-18 11:10:27 +00:00
def value_forward(self):
pass
2019-09-21 15:17:09 +00:00
2019-09-18 11:10:27 +00:00
MlpPolicy = PPOPolicy
register_policy("MlpPolicy", MlpPolicy)