stable-baselines3/torchy_baselines/ppo/policies.py

144 lines
6.4 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-11-22 12:06:41 +00:00
from torchy_baselines.common.policies import BasePolicy, register_policy, MlpExtractor
2019-10-28 17:24:13 +00:00
from torchy_baselines.common.distributions import make_proba_distribution,\
DiagGaussianDistribution, CategoricalDistribution, StateDependentNoiseDistribution
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):
2019-11-22 16:24:47 +00:00
"""
Policy class (with both actor and critic) for A2C and derivates (PPO).
:param observation_space: (gym.spaces.Space) Observation space
:param action_dim: (gym.spaces.Space) Action space
:param learning_rate: (callable) Learning rate schedule (could be constant)
:param net_arch: ([int or dict]) The specification of the policy and value networks.
:param device: (str or th.device) Device on which the code should run.
:param activation_fn: (nn.Module) Activation function
:param adam_epsilon: (float) Small values to avoid NaN in ADAM optimizer
:param ortho_init: (bool) Whether to use or not orthogonal initialization
:param use_sde: (bool) Whether to use State Dependent Exploration or not
:param log_std_init: (float) Initial value for the log standard deviation
"""
2019-09-18 11:10:27 +00:00
def __init__(self, observation_space, action_space,
2019-10-28 15:47:13 +00:00
learning_rate, net_arch=None, device='cpu',
2019-10-28 17:24:13 +00:00
activation_fn=nn.Tanh, adam_epsilon=1e-5,
2019-10-29 17:43:16 +00:00
ortho_init=True, use_sde=False, log_std_init=0.0):
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-11-22 16:24:14 +00:00
# Default network architecture, from stable-baselines
2019-09-18 11:10:27 +00:00
if net_arch is None:
2019-11-22 16:24:14 +00:00
net_arch = [dict(pi=[64, 64], vf=[64, 64])]
2019-09-18 11:10:27 +00:00
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-10-17 11:32:25 +00:00
# In the future, feature_extractor will be replaced with a CNN
self.features_extractor = nn.Flatten()
self.features_dim = self.obs_dim
2019-10-29 17:43:16 +00:00
self.log_std_init = log_std_init
2019-10-28 17:24:13 +00:00
# Action distribution
2019-10-31 10:44:27 +00:00
self.action_dist = make_proba_distribution(action_space, use_sde=use_sde)
2019-10-28 17:24:13 +00:00
2019-09-18 11:10:27 +00:00
self._build(learning_rate)
2019-10-28 17:24:13 +00:00
def reset_noise_net(self):
2019-11-25 12:19:33 +00:00
"""
Sample new weights for the exploration matrix.
"""
2019-10-28 17:24:13 +00:00
self.action_dist.sample_weights(self.log_std)
2019-09-18 11:10:27 +00:00
def _build(self, learning_rate):
2019-10-17 11:32:25 +00:00
self.mlp_extractor = MlpExtractor(self.features_dim, net_arch=self.net_arch,
activation_fn=self.activation_fn, device=self.device)
2019-09-21 16:12:06 +00:00
2019-10-28 17:24:13 +00:00
if isinstance(self.action_dist, (DiagGaussianDistribution, StateDependentNoiseDistribution)):
2019-10-29 17:43:16 +00:00
self.action_net, self.log_std = self.action_dist.proba_distribution_net(latent_dim=self.mlp_extractor.latent_dim_pi,
log_std_init=self.log_std_init)
elif isinstance(self.action_dist, CategoricalDistribution):
2019-10-17 11:32:25 +00:00
self.action_net = self.action_dist.proba_distribution_net(latent_dim=self.mlp_extractor.latent_dim_pi)
2019-10-17 11:32:25 +00:00
self.value_net = nn.Linear(self.mlp_extractor.latent_dim_vf, 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:
2019-10-17 11:32:25 +00:00
for module in [self.mlp_extractor, self.action_net, self.value_net]:
2019-11-25 12:19:33 +00:00
# Values from stable-baselines, TODO: check why
2019-09-26 14:29:47 +00:00
gain = {
2019-10-17 11:32:25 +00:00
self.mlp_extractor: np.sqrt(2),
2019-09-26 14:29:47 +00:00
self.action_net: 0.01,
self.value_net: 1
}[module]
module.apply(partial(self.init_weights, gain=gain))
2019-10-28 15:47:13 +00:00
self.optimizer = th.optim.Adam(self.parameters(), lr=learning_rate(1), eps=self.adam_epsilon)
2019-09-18 11:10:27 +00:00
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)
2019-10-31 10:44:27 +00:00
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)
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-10-17 11:32:25 +00:00
return self.mlp_extractor(self.features_extractor(obs))
2019-09-21 16:12:06 +00:00
2019-10-31 10:44:27 +00:00
def _get_action_dist_from_latent(self, latent_pi, deterministic=False):
mean_actions = self.action_net(latent_pi)
2019-10-29 17:43:16 +00:00
if isinstance(self.action_dist, DiagGaussianDistribution):
return self.action_dist.proba_distribution(mean_actions, self.log_std, deterministic=deterministic)
2019-10-29 17:43:16 +00:00
elif isinstance(self.action_dist, CategoricalDistribution):
return self.action_dist.proba_distribution(mean_actions, deterministic=deterministic)
2019-10-29 17:43:16 +00:00
2019-10-28 17:24:13 +00:00
elif isinstance(self.action_dist, StateDependentNoiseDistribution):
2019-10-31 10:44:27 +00:00
return self.action_dist.proba_distribution(mean_actions, self.log_std, latent_pi, 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-10-31 10:44:27 +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-11-22 10:42:58 +00:00
def evaluate_actions(self, obs, action, deterministic=False):
2019-11-22 16:24:47 +00:00
"""
Evaluate actions according to the current policy,
given the observations.
:param obs: (th.Tensor)
:param action: (th.Tensor)
:param deterministic: (bool)
:return: (th.Tensor, th.Tensor, th.Tensor) estimated value, log likelihood of taking those actions
and entropy of the action distribution.
"""
2019-09-24 12:53:03 +00:00
latent_pi, latent_vf = self._get_latent(obs)
2019-10-31 10:44:27 +00:00
_, 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)
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-11-22 16:24:47 +00:00
def value_forward(self, obs):
_, latent_vf = self._get_latent(obs)
return self.value_net(latent_vf)
2019-09-18 11:10:27 +00:00
2019-09-21 15:17:09 +00:00
2019-09-18 11:10:27 +00:00
MlpPolicy = PPOPolicy
register_policy("MlpPolicy", MlpPolicy)