Smaller conv nets

This commit is contained in:
ebuehrle
2022-01-21 10:23:59 +01:00
parent 52d4b5a36c
commit e1f2d58255
2 changed files with 51 additions and 8 deletions

View File

@@ -14,16 +14,14 @@ class CnnDiscriminator(torch.nn.Module):
in_channels = obs_channels + action_size
self.cnn = torch.nn.Sequential(
torch.nn.Conv2d(in_channels, 32, kernel_size=(8, 8), stride=(4, 4)), # 5+1 -> 32
torch.nn.Conv2d(in_channels, 4, kernel_size=8, stride=4, padding=0),
torch.nn.ReLU(),
torch.nn.Conv2d(32, 64, kernel_size=(4, 4), stride=(2, 2)), # 32 -> 64
torch.nn.Conv2d(4, 8, kernel_size=4, stride=2, padding=0),
torch.nn.ReLU(),
torch.nn.Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1)), # 64 -> 64
torch.nn.Flatten(start_dim=-3, end_dim=-1),
torch.nn.LazyLinear(512),
torch.nn.ReLU(),
torch.nn.Flatten(start_dim=1, end_dim=-1),
torch.nn.LazyLinear(512), # 28224 -> 512
torch.nn.ReLU(),
torch.nn.LazyLinear(1), # 512 -> 1
torch.nn.LazyLinear(1),
)
@staticmethod

View File

@@ -1,13 +1,58 @@
from stable_baselines3.common.policies import ActorCriticPolicy, ActorCriticCnnPolicy
from torch.distributions import Categorical
import gym
import torch
import torch.nn as nn
from stable_baselines3.common.torch_layers import BaseFeaturesExtractor
from stable_baselines3.common.preprocessing import is_image_space
class CustomCNN(BaseFeaturesExtractor):
"""
Smaller version of `stable_baselines3.common.torch_layers.NatureCNN`
:param observation_space:
:param features_dim: Number of features extracted.
This corresponds to the number of unit for the last layer.
"""
def __init__(self, observation_space: gym.spaces.Box, features_dim: int = 512):
super().__init__(observation_space, features_dim)
# We assume CxHxW images (channels first)
# Re-ordering will be done by pre-preprocessing or wrapper
assert is_image_space(observation_space, check_channels=False), (
"You should use CustomCNN "
f"only with images not with {observation_space}\n"
"(you are probably using `CnnPolicy` instead of `MlpPolicy` or `MultiInputPolicy`)\n"
"If you are using a custom environment,\n"
"please check it using our env checker:\n"
"https://stable-baselines3.readthedocs.io/en/master/common/env_checker.html"
)
n_input_channels = observation_space.shape[0]
self.cnn = nn.Sequential(
nn.Conv2d(n_input_channels, 4, kernel_size=8, stride=4, padding=0),
nn.ReLU(),
nn.Conv2d(4, 8, kernel_size=4, stride=2, padding=0),
nn.ReLU(),
nn.Flatten(),
)
# Compute shape by doing one forward pass
with torch.no_grad():
n_flatten = self.cnn(torch.as_tensor(observation_space.sample()[None]).float()).shape[1]
self.linear = nn.Sequential(nn.Linear(n_flatten, features_dim), nn.ReLU())
def forward(self, observations: torch.Tensor) -> torch.Tensor:
return self.linear(self.cnn(observations))
class OptionsCnnPolicy(ActorCriticPolicy):
"""
Class for high-level options policy (generator)
"""
def __init__(self, observation_space, *args, eps=0, **kwargs):
super().__init__(observation_space, *args, **kwargs)
self.cnn_policy = ActorCriticCnnPolicy(observation_space['obs'], *args, **kwargs)
self.cnn_policy = ActorCriticCnnPolicy(observation_space['obs'], *args, features_extractor_class=CustomCNN, **kwargs)
self.eps = eps
def _prior_distribution(self, s):