mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-04 20:23:54 +00:00
190 lines
9.4 KiB
Python
190 lines
9.4 KiB
Python
from functools import partial
|
|
from itertools import zip_longest
|
|
|
|
import torch as th
|
|
import torch.nn as nn
|
|
import numpy as np
|
|
|
|
from torchy_baselines.common.policies import BasePolicy, register_policy, create_mlp
|
|
from torchy_baselines.common.distributions import make_proba_distribution, DiagGaussianDistribution, CategoricalDistribution
|
|
|
|
|
|
class MlpExtractor(nn.Module):
|
|
"""
|
|
Constructs an MLP that receives observations as an input and outputs a latent representation for the policy and
|
|
a value network. The ``net_arch`` parameter allows to specify the amount and size of the hidden layers and how many
|
|
of them are shared between the policy network and the value network. It is assumed to be a list with the following
|
|
structure:
|
|
|
|
1. An arbitrary length (zero allowed) number of integers each specifying the number of units in a shared layer.
|
|
If the number of ints is zero, there will be no shared layers.
|
|
2. An optional dict, to specify the following non-shared layers for the value network and the policy network.
|
|
It is formatted like ``dict(vf=[<value layer sizes>], pi=[<policy layer sizes>])``.
|
|
If it is missing any of the keys (pi or vf), no non-shared layers (empty list) is assumed.
|
|
|
|
For example to construct a network with one shared layer of size 55 followed by two non-shared layers for the value
|
|
network of size 255 and a single non-shared layer of size 128 for the policy network, the following layers_spec
|
|
would be used: ``[55, dict(vf=[255, 255], pi=[128])]``. A simple shared network topology with two layers of size 128
|
|
would be specified as [128, 128].
|
|
|
|
Adapted from Stable Baselines.
|
|
|
|
:param flat_observations: (th.Tensor) The observations to base policy and value function on.
|
|
:param net_arch: ([int or dict]) The specification of the policy and value networks.
|
|
See above for details on its formatting.
|
|
:param activation_fn: (nn.Module) The activation function to use for the networks.
|
|
:param device: (th.device)
|
|
"""
|
|
def __init__(self, feature_dim, net_arch, activation_fn, device='cpu'):
|
|
super(MlpExtractor, self).__init__()
|
|
|
|
shared_net, policy_net, value_net = [], [], []
|
|
policy_only_layers = [] # Layer sizes of the network that only belongs to the policy network
|
|
value_only_layers = [] # Layer sizes of the network that only belongs to the value network
|
|
last_layer_dim_shared = feature_dim
|
|
|
|
# Iterate through the shared layers and build the shared parts of the network
|
|
for idx, layer in enumerate(net_arch):
|
|
if isinstance(layer, int): # Check that this is a shared layer
|
|
layer_size = layer
|
|
# TODO: give layer a meaningful name
|
|
shared_net.append(nn.Linear(last_layer_dim_shared, layer_size))
|
|
shared_net.append(activation_fn())
|
|
last_layer_dim_shared = layer_size
|
|
else:
|
|
assert isinstance(layer, dict), "Error: the net_arch list can only contain ints and dicts"
|
|
if 'pi' in layer:
|
|
assert isinstance(layer['pi'], list), "Error: net_arch[-1]['pi'] must contain a list of integers."
|
|
policy_only_layers = layer['pi']
|
|
|
|
if 'vf' in layer:
|
|
assert isinstance(layer['vf'], list), "Error: net_arch[-1]['vf'] must contain a list of integers."
|
|
value_only_layers = layer['vf']
|
|
break # From here on the network splits up in policy and value network
|
|
|
|
last_layer_dim_pi = last_layer_dim_shared
|
|
last_layer_dim_vf = last_layer_dim_shared
|
|
|
|
# Build the non-shared part of the network
|
|
for idx, (pi_layer_size, vf_layer_size) in enumerate(zip_longest(policy_only_layers, value_only_layers)):
|
|
if pi_layer_size is not None:
|
|
assert isinstance(pi_layer_size, int), "Error: net_arch[-1]['pi'] must only contain integers."
|
|
policy_net.append(nn.Linear(last_layer_dim_pi, pi_layer_size))
|
|
policy_net.append(activation_fn())
|
|
last_layer_dim_pi = pi_layer_size
|
|
|
|
if vf_layer_size is not None:
|
|
assert isinstance(vf_layer_size, int), "Error: net_arch[-1]['vf'] must only contain integers."
|
|
value_net.append(nn.Linear(last_layer_dim_vf, vf_layer_size))
|
|
value_net.append(activation_fn())
|
|
last_layer_dim_vf = vf_layer_size
|
|
|
|
# Save dim, used to create the distributions
|
|
self.latent_dim_pi = last_layer_dim_pi
|
|
self.latent_dim_vf = last_layer_dim_vf
|
|
|
|
# Create networks
|
|
# If the list of layers is empty, the network will just act as an Identity module
|
|
self.shared_net = nn.Sequential(*shared_net).to(device)
|
|
self.policy_net = nn.Sequential(*policy_net).to(device)
|
|
self.value_net = nn.Sequential(*value_net).to(device)
|
|
|
|
def forward(self, features):
|
|
"""
|
|
:return: (th.Tensor, th.Tensor) latent_policy, latent_value of the specified network.
|
|
If all layers are shared, then ``latent_policy == latent_value``
|
|
"""
|
|
shared_latent = self.shared_net(features)
|
|
return self.policy_net(shared_latent), self.value_net(shared_latent)
|
|
|
|
|
|
class PPOPolicy(BasePolicy):
|
|
def __init__(self, observation_space, action_space,
|
|
learning_rate, net_arch=None, device='cpu',
|
|
activation_fn=nn.Tanh, adam_epsilon=1e-5, ortho_init=True):
|
|
super(PPOPolicy, self).__init__(observation_space, action_space, device)
|
|
self.obs_dim = self.observation_space.shape[0]
|
|
if net_arch is None:
|
|
net_arch = [dict(pi=[64], vf=[64])]
|
|
self.net_arch = net_arch
|
|
self.activation_fn = activation_fn
|
|
self.adam_epsilon = adam_epsilon
|
|
self.ortho_init = ortho_init
|
|
self.net_args = {
|
|
'input_dim': self.obs_dim,
|
|
'output_dim': -1,
|
|
'net_arch': self.net_arch,
|
|
'activation_fn': self.activation_fn
|
|
}
|
|
self.shared_net = None
|
|
self.pi_net, self.vf_net = None, None
|
|
# Action distribution
|
|
self.action_dist = make_proba_distribution(action_space)
|
|
# In the future, feature_extractor will be replaced with a CNN
|
|
self.features_extractor = nn.Flatten()
|
|
self.features_dim = self.obs_dim
|
|
self._build(learning_rate)
|
|
|
|
def _build(self, learning_rate):
|
|
self.mlp_extractor = MlpExtractor(self.features_dim, net_arch=self.net_arch,
|
|
activation_fn=self.activation_fn, device=self.device)
|
|
|
|
# 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.mlp_extractor.latent_dim_pi)
|
|
elif isinstance(self.action_dist, CategoricalDistribution):
|
|
self.action_net = self.action_dist.proba_distribution_net(latent_dim=self.mlp_extractor.latent_dim_pi)
|
|
|
|
self.value_net = nn.Linear(self.mlp_extractor.latent_dim_vf, 1)
|
|
# Init weights: use orthogonal initialization
|
|
# with small initial weight for the output
|
|
if self.ortho_init:
|
|
for module in [self.mlp_extractor, self.action_net, self.value_net]:
|
|
# Values from stable-baselines check why
|
|
gain = {
|
|
self.mlp_extractor: np.sqrt(2),
|
|
self.action_net: 0.01,
|
|
self.value_net: 1
|
|
}[module]
|
|
module.apply(partial(self.init_weights, gain=gain))
|
|
self.optimizer = th.optim.Adam(self.parameters(), lr=learning_rate(1), eps=self.adam_epsilon)
|
|
|
|
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)
|
|
value = self.value_net(latent_vf)
|
|
action, action_distribution = self._get_action_dist_from_latent(latent_pi, deterministic=deterministic)
|
|
log_prob = action_distribution.log_prob(action)
|
|
return action, value, log_prob
|
|
|
|
def _get_latent(self, obs):
|
|
return self.mlp_extractor(self.features_extractor(obs))
|
|
|
|
def _get_action_dist_from_latent(self, latent, deterministic=False):
|
|
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)
|
|
|
|
def actor_forward(self, obs, deterministic=False):
|
|
latent_pi, _ = self._get_latent(obs)
|
|
action, _ = self._get_action_dist_from_latent(latent_pi, deterministic=deterministic)
|
|
return action.detach().cpu().numpy()
|
|
|
|
def get_policy_stats(self, obs, action):
|
|
latent_pi, latent_vf = self._get_latent(obs)
|
|
_, action_distribution = self._get_action_dist_from_latent(latent_pi)
|
|
log_prob = action_distribution.log_prob(action)
|
|
value = self.value_net(latent_vf)
|
|
return value, log_prob, action_distribution.entropy()
|
|
|
|
def value_forward(self):
|
|
pass
|
|
|
|
|
|
MlpPolicy = PPOPolicy
|
|
|
|
register_policy("MlpPolicy", MlpPolicy)
|