Remove default seed and bump dependencies

This commit is contained in:
Antonin Raffin 2020-03-10 17:43:54 +01:00
parent 80fb62e22d
commit 6ebad92e1b
8 changed files with 76 additions and 42 deletions

View file

@ -9,6 +9,8 @@ Pre-Release 0.3.0a0 (WIP)
Breaking Changes:
^^^^^^^^^^^^^^^^^
- Removed default seed
- Bump dependencies (PyTorch and Gym)
New Features:
^^^^^^^^^^^^^
@ -24,6 +26,7 @@ Others:
- SAC with SDE now sample only one matrix
- Added ``clip_mean`` parameter to SAC policy
- Buffers now return ``NamedTuple``
- More typing
Documentation:
^^^^^^^^^^^^^^

View file

@ -7,9 +7,9 @@ setup(name='torchy_baselines',
packages=[package for package in find_packages()
if package.startswith('torchy_baselines')],
install_requires=[
'gym[classic_control]>=0.10.9',
'gym[classic_control]>=0.11',
'numpy',
'torch>=1.2.0',
'torch>=1.4.0',
'cloudpickle',
# For reading logs
'pandas',

View file

@ -51,7 +51,7 @@ class A2C(PPO):
ent_coef=0.0, vf_coef=0.5, max_grad_norm=0.5,
rms_prop_eps=1e-5, use_rms_prop=True, use_sde=False, sde_sample_freq=-1,
normalize_advantage=False, tensorboard_log=None, create_eval_env=False,
policy_kwargs=None, verbose=0, seed=0, device='auto',
policy_kwargs=None, verbose=0, seed=None, device='auto',
_init_setup_model=True):
super(A2C, self).__init__(policy, env, learning_rate=learning_rate,

View file

@ -62,7 +62,7 @@ class CEMRL(TD3):
action_noise=None, target_policy_noise=0.2, target_noise_clip=0.5,
n_episodes_rollout=1, update_style='original',
tensorboard_log=None, create_eval_env=False,
policy_kwargs=None, verbose=0, seed=0, device='auto',
policy_kwargs=None, verbose=0, seed=None, device='auto',
_init_setup_model=True):
super(CEMRL, self).__init__(policy, env,

View file

@ -79,7 +79,7 @@ class PPO(BaseRLModel):
ent_coef=0.0, vf_coef=0.5, max_grad_norm=0.5,
use_sde=False, sde_sample_freq=-1,
target_kl=None, tensorboard_log=None, create_eval_env=False,
policy_kwargs=None, verbose=0, seed=0, device='auto',
policy_kwargs=None, verbose=0, seed=None, device='auto',
_init_setup_model=True):
super(PPO, self).__init__(policy, env, PPOPolicy, policy_kwargs=policy_kwargs,

View file

@ -1,5 +1,6 @@
from typing import Optional, List, Tuple
from typing import Optional, List, Tuple, Callable, Union
import gym
import torch as th
import torch.nn as nn
@ -143,8 +144,10 @@ class Critic(BaseNetwork):
:param net_arch: ([int]) Network architecture
:param activation_fn: (nn.Module) Activation function
"""
def __init__(self, obs_dim, action_dim,
net_arch, activation_fn=nn.ReLU):
def __init__(self, obs_dim: int,
action_dim: int,
net_arch: List[int],
activation_fn: nn.Module = nn.ReLU):
super(Critic, self).__init__()
q1_net = create_mlp(obs_dim + action_dim, 1, net_arch, activation_fn)
@ -155,13 +158,10 @@ class Critic(BaseNetwork):
self.q_networks = [self.q1_net, self.q2_net]
def forward(self, obs, action):
def forward(self, obs: th.Tensor, action: th.Tensor) -> List[th.Tensor]:
qvalue_input = th.cat([obs, action], dim=1)
return [q_net(qvalue_input) for q_net in self.q_networks]
def q1_forward(self, obs, action):
return self.q_networks[0](th.cat([obs, action], dim=1))
class SACPolicy(BasePolicy):
"""
@ -183,11 +183,17 @@ class SACPolicy(BasePolicy):
above zero and prevent it from growing too fast. In practice, `exp()` is usually enough.
:param clip_mean: (float) Clip the mean output when using SDE to avoid numerical instability.
"""
def __init__(self, observation_space, action_space,
learning_rate, net_arch=None, device='cpu',
activation_fn=nn.ReLU, use_sde=False,
log_std_init=-3, sde_net_arch=None,
use_expln=False, clip_mean=2.0):
def __init__(self, observation_space: gym.spaces.Space,
action_space: gym.spaces.Space,
learning_rate: Callable,
net_arch: Optional[List[int]] = None,
device: Union[th.device, str] = 'cpu',
activation_fn: nn.Module = nn.ReLU,
use_sde: bool = False,
log_std_init: float = -3,
sde_net_arch: Optional[List[int]] = None,
use_expln: bool = False,
clip_mean: float = 2.0):
super(SACPolicy, self).__init__(observation_space, action_space, device, squash_output=True)
if net_arch is None:
@ -217,7 +223,7 @@ class SACPolicy(BasePolicy):
self._build(learning_rate)
def _build(self, learning_rate):
def _build(self, learning_rate: Callable) -> None:
self.actor = self.make_actor()
self.actor.optimizer = th.optim.Adam(self.actor.parameters(), lr=learning_rate(1))
@ -226,13 +232,13 @@ class SACPolicy(BasePolicy):
self.critic_target.load_state_dict(self.critic.state_dict())
self.critic.optimizer = th.optim.Adam(self.critic.parameters(), lr=learning_rate(1))
def make_actor(self):
def make_actor(self) -> Actor:
return Actor(**self.actor_kwargs).to(self.device)
def make_critic(self):
def make_critic(self) -> Critic:
return Critic(**self.net_args).to(self.device)
def forward(self, obs):
def forward(self, obs: th.Tensor) -> th.Tensor:
return self.actor(obs)
def predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:

View file

@ -1,13 +1,16 @@
from typing import List, Tuple
from typing import List, Tuple, Type, Union, Callable, Optional, Dict, Any
import torch as th
import torch.nn.functional as F
import numpy as np
from torchy_baselines.common import logger
from torchy_baselines.common.base_class import OffPolicyRLModel
from torchy_baselines.common.buffers import ReplayBuffer
from torchy_baselines.common.type_aliases import GymEnv
from torchy_baselines.common.noise import ActionNoise
from torchy_baselines.common.callbacks import BaseCallback
from torchy_baselines.sac.policies import SACPolicy
from torchy_baselines.common import logger
class SAC(OffPolicyRLModel):
@ -25,7 +28,7 @@ class SAC(OffPolicyRLModel):
in https://github.com/hill-a/stable-baselines/issues/270
:param policy: (SACPolicy or str) The policy model to use (MlpPolicy, CnnPolicy, ...)
:param env: (Gym environment or str) The environment to learn from (if registered in Gym, can be str)
:param env: (GymEnv or str) The environment to learn from (if registered in Gym, can be str)
:param learning_rate: (float or callable) learning rate for adam optimizer,
the same learning rate will be used for all networks (Q-Values, Actor and Value function)
it can be a function of the current progress (from 1 to 0)
@ -61,16 +64,31 @@ class SAC(OffPolicyRLModel):
:param _init_setup_model: (bool) Whether or not to build the network at the creation of the instance
"""
def __init__(self, policy, env, learning_rate=3e-4, buffer_size=int(1e6),
learning_starts=100, batch_size=256,
tau=0.005, ent_coef='auto', target_update_interval=1,
train_freq=1, gradient_steps=1, n_episodes_rollout=-1,
target_entropy='auto', action_noise=None,
gamma=0.99, use_sde=False, sde_sample_freq=-1,
use_sde_at_warmup=False,
tensorboard_log=None, create_eval_env=False,
policy_kwargs=None, verbose=0, seed=0, device='auto',
_init_setup_model=True):
def __init__(self, policy: Union[str, Type[SACPolicy]],
env: Union[GymEnv, str],
learning_rate: Union[float, Callable] = 3e-4,
buffer_size: int = int(1e6),
learning_starts: int = 100,
batch_size: int = 256,
tau: float = 0.005,
ent_coef: Union[str, float] = 'auto',
target_update_interval: int = 1,
train_freq: int = 1,
gradient_steps: int = 1,
n_episodes_rollout: int = -1,
target_entropy: Union[str, float] = 'auto',
action_noise: Optional[ActionNoise] = None,
gamma: float = 0.99,
use_sde: bool = False,
sde_sample_freq: int = -1,
use_sde_at_warmup: bool = False,
tensorboard_log: Optional[str] = None,
create_eval_env: bool = False,
policy_kwargs: Dict[str, Any] = None,
verbose: int = 0,
seed: Optional[int] = None,
device: Union[th.device, str] = 'auto',
_init_setup_model: bool = True):
super(SAC, self).__init__(policy, env, SACPolicy, policy_kwargs, verbose, device,
create_eval_env=create_eval_env, seed=seed,
@ -79,7 +97,7 @@ class SAC(OffPolicyRLModel):
self.learning_rate = learning_rate
self.target_entropy = target_entropy
self.log_ent_coef = None
self.log_ent_coef = None # type: Optional[th.Tensor]
self.target_update_interval = target_update_interval
self.buffer_size = buffer_size
# In the original paper, same learning rate is used for all networks
@ -101,7 +119,7 @@ class SAC(OffPolicyRLModel):
if _init_setup_model:
self._setup_model()
def _setup_model(self):
def _setup_model(self) -> None:
self._setup_learning_rate()
obs_dim, action_dim = self.observation_space.shape[0], self.action_space.shape[0]
if self.seed is not None:
@ -143,12 +161,12 @@ class SAC(OffPolicyRLModel):
self.policy = self.policy.to(self.device)
self._create_aliases()
def _create_aliases(self):
def _create_aliases(self) -> None:
self.actor = self.policy.actor
self.critic = self.policy.critic
self.critic_target = self.policy.critic_target
def train(self, gradient_steps: int, batch_size: int = 64):
def train(self, gradient_steps: int, batch_size: int = 64) -> None:
# Update optimizers learning rate
optimizers = [self.actor.optimizer, self.critic.optimizer]
if self.ent_coef_optimizer is not None:
@ -233,9 +251,16 @@ class SAC(OffPolicyRLModel):
if ent_coef_loss is not None:
logger.logkv("ent_coef_loss", ent_coef_loss.item())
def learn(self, total_timesteps, callback=None, log_interval=4,
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="SAC",
eval_log_path=None, reset_num_timesteps=True):
def learn(self,
total_timesteps: int,
callback: Optional[BaseCallback] = None,
log_interval: int = 4,
eval_env: Optional[GymEnv] = None,
eval_freq: int = -1,
n_eval_episodes: int = 5,
tb_log_name: str = "SAC",
eval_log_path: Optional[str] = None,
reset_num_timesteps: bool = True) -> OffPolicyRLModel:
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq,
n_eval_episodes, eval_log_path, reset_num_timesteps)

View file

@ -65,7 +65,7 @@ class TD3(OffPolicyRLModel):
use_sde=False, sde_sample_freq=-1, sde_max_grad_norm=1,
sde_ent_coef=0.0, sde_log_std_scheduler=None, use_sde_at_warmup=False,
tensorboard_log=None, create_eval_env=False, policy_kwargs=None, verbose=0,
seed=0, device='auto', _init_setup_model=True):
seed=None, device='auto', _init_setup_model=True):
super(TD3, self).__init__(policy, env, TD3Policy, policy_kwargs, verbose, device,
create_eval_env=create_eval_env, seed=seed,