mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-17 22:30:59 +00:00
Add supported action spaces checks (#254)
* Add supported action spaces checks * Address comment
This commit is contained in:
parent
e747e7e2b3
commit
2b9fc1f923
10 changed files with 53 additions and 3 deletions
|
|
@ -42,6 +42,7 @@ Others:
|
||||||
- Add more issue templates
|
- Add more issue templates
|
||||||
- Add signatures to callable type annotations (@erniejunior)
|
- Add signatures to callable type annotations (@erniejunior)
|
||||||
- Improve error message in ``NatureCNN``
|
- Improve error message in ``NatureCNN``
|
||||||
|
- Added checks for supported action spaces to improve clarity of error messages for the user
|
||||||
|
|
||||||
Documentation:
|
Documentation:
|
||||||
^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^
|
||||||
|
|
|
||||||
|
|
@ -96,6 +96,12 @@ class A2C(OnPolicyAlgorithm):
|
||||||
create_eval_env=create_eval_env,
|
create_eval_env=create_eval_env,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
_init_setup_model=False,
|
_init_setup_model=False,
|
||||||
|
supported_action_spaces=(
|
||||||
|
spaces.Box,
|
||||||
|
spaces.Discrete,
|
||||||
|
spaces.MultiDiscrete,
|
||||||
|
spaces.MultiBinary,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.normalize_advantage = normalize_advantage
|
self.normalize_advantage = normalize_advantage
|
||||||
|
|
|
||||||
|
|
@ -79,6 +79,7 @@ class BaseAlgorithm(ABC):
|
||||||
instead of action noise exploration (default: False)
|
instead of action noise exploration (default: False)
|
||||||
:param sde_sample_freq: Sample a new noise matrix every n steps when using gSDE
|
:param sde_sample_freq: Sample a new noise matrix every n steps when using gSDE
|
||||||
Default: -1 (only sample at the beginning of the rollout)
|
Default: -1 (only sample at the beginning of the rollout)
|
||||||
|
:param supported_action_spaces: The action spaces supported by the algorithm.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
|
|
@ -97,6 +98,7 @@ class BaseAlgorithm(ABC):
|
||||||
seed: Optional[int] = None,
|
seed: Optional[int] = None,
|
||||||
use_sde: bool = False,
|
use_sde: bool = False,
|
||||||
sde_sample_freq: int = -1,
|
sde_sample_freq: int = -1,
|
||||||
|
supported_action_spaces: Optional[Tuple[gym.spaces.Space, ...]] = None,
|
||||||
):
|
):
|
||||||
|
|
||||||
if isinstance(policy, str) and policy_base is not None:
|
if isinstance(policy, str) and policy_base is not None:
|
||||||
|
|
@ -158,6 +160,12 @@ class BaseAlgorithm(ABC):
|
||||||
self.n_envs = env.num_envs
|
self.n_envs = env.num_envs
|
||||||
self.env = env
|
self.env = env
|
||||||
|
|
||||||
|
if supported_action_spaces is not None:
|
||||||
|
assert isinstance(self.action_space, supported_action_spaces), (
|
||||||
|
f"The algorithm only supports {supported_action_spaces} as action spaces "
|
||||||
|
f"but {self.action_space} was provided"
|
||||||
|
)
|
||||||
|
|
||||||
if not support_multi_env and self.n_envs > 1:
|
if not support_multi_env and self.n_envs > 1:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Error: the model does not support multiple envs; it requires " "a single vectorized environment."
|
"Error: the model does not support multiple envs; it requires " "a single vectorized environment."
|
||||||
|
|
|
||||||
|
|
@ -69,6 +69,7 @@ class OffPolicyAlgorithm(BaseAlgorithm):
|
||||||
:param sde_support: Whether the model support gSDE or not
|
:param sde_support: Whether the model support gSDE or not
|
||||||
:param remove_time_limit_termination: Remove terminations (dones) that are due to time limit.
|
:param remove_time_limit_termination: Remove terminations (dones) that are due to time limit.
|
||||||
See https://github.com/hill-a/stable-baselines/issues/863
|
See https://github.com/hill-a/stable-baselines/issues/863
|
||||||
|
:param supported_action_spaces: The action spaces supported by the algorithm.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
|
|
@ -100,6 +101,7 @@ class OffPolicyAlgorithm(BaseAlgorithm):
|
||||||
use_sde_at_warmup: bool = False,
|
use_sde_at_warmup: bool = False,
|
||||||
sde_support: bool = True,
|
sde_support: bool = True,
|
||||||
remove_time_limit_termination: bool = False,
|
remove_time_limit_termination: bool = False,
|
||||||
|
supported_action_spaces: Optional[Tuple[gym.spaces.Space, ...]] = None,
|
||||||
):
|
):
|
||||||
|
|
||||||
super(OffPolicyAlgorithm, self).__init__(
|
super(OffPolicyAlgorithm, self).__init__(
|
||||||
|
|
@ -117,6 +119,7 @@ class OffPolicyAlgorithm(BaseAlgorithm):
|
||||||
seed=seed,
|
seed=seed,
|
||||||
use_sde=use_sde,
|
use_sde=use_sde,
|
||||||
sde_sample_freq=sde_sample_freq,
|
sde_sample_freq=sde_sample_freq,
|
||||||
|
supported_action_spaces=supported_action_spaces,
|
||||||
)
|
)
|
||||||
self.buffer_size = buffer_size
|
self.buffer_size = buffer_size
|
||||||
self.batch_size = batch_size
|
self.batch_size = batch_size
|
||||||
|
|
|
||||||
|
|
@ -46,6 +46,7 @@ class OnPolicyAlgorithm(BaseAlgorithm):
|
||||||
:param device: Device (cpu, cuda, ...) on which the code should be run.
|
:param device: Device (cpu, cuda, ...) on which the code should be run.
|
||||||
Setting it to auto, the code will be run on the GPU if possible.
|
Setting it to auto, the code will be run on the GPU if possible.
|
||||||
:param _init_setup_model: Whether or not to build the network at the creation of the instance
|
:param _init_setup_model: Whether or not to build the network at the creation of the instance
|
||||||
|
:param supported_action_spaces: The action spaces supported by the algorithm.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
|
|
@ -69,6 +70,7 @@ class OnPolicyAlgorithm(BaseAlgorithm):
|
||||||
seed: Optional[int] = None,
|
seed: Optional[int] = None,
|
||||||
device: Union[th.device, str] = "auto",
|
device: Union[th.device, str] = "auto",
|
||||||
_init_setup_model: bool = True,
|
_init_setup_model: bool = True,
|
||||||
|
supported_action_spaces: Optional[Tuple[gym.spaces.Space, ...]] = None,
|
||||||
):
|
):
|
||||||
|
|
||||||
super(OnPolicyAlgorithm, self).__init__(
|
super(OnPolicyAlgorithm, self).__init__(
|
||||||
|
|
@ -85,6 +87,7 @@ class OnPolicyAlgorithm(BaseAlgorithm):
|
||||||
support_multi_env=True,
|
support_multi_env=True,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
tensorboard_log=tensorboard_log,
|
tensorboard_log=tensorboard_log,
|
||||||
|
supported_action_spaces=supported_action_spaces,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.n_steps = n_steps
|
self.n_steps = n_steps
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
from typing import Any, Dict, List, Optional, Tuple, Type, Union
|
from typing import Any, Dict, List, Optional, Tuple, Type, Union
|
||||||
|
|
||||||
|
import gym
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch as th
|
import torch as th
|
||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
|
|
@ -105,6 +106,7 @@ class DQN(OffPolicyAlgorithm):
|
||||||
seed=seed,
|
seed=seed,
|
||||||
sde_support=False,
|
sde_support=False,
|
||||||
optimize_memory_usage=optimize_memory_usage,
|
optimize_memory_usage=optimize_memory_usage,
|
||||||
|
supported_action_spaces=(gym.spaces.Discrete,),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.exploration_initial_eps = exploration_initial_eps
|
self.exploration_initial_eps = exploration_initial_eps
|
||||||
|
|
|
||||||
|
|
@ -108,6 +108,12 @@ class PPO(OnPolicyAlgorithm):
|
||||||
create_eval_env=create_eval_env,
|
create_eval_env=create_eval_env,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
_init_setup_model=False,
|
_init_setup_model=False,
|
||||||
|
supported_action_spaces=(
|
||||||
|
spaces.Box,
|
||||||
|
spaces.Discrete,
|
||||||
|
spaces.MultiDiscrete,
|
||||||
|
spaces.MultiBinary,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.batch_size = batch_size
|
self.batch_size = batch_size
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
from typing import Any, Dict, List, Optional, Tuple, Type, Union
|
from typing import Any, Dict, List, Optional, Tuple, Type, Union
|
||||||
|
|
||||||
|
import gym
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch as th
|
import torch as th
|
||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
|
|
@ -124,6 +125,7 @@ class SAC(OffPolicyAlgorithm):
|
||||||
sde_sample_freq=sde_sample_freq,
|
sde_sample_freq=sde_sample_freq,
|
||||||
use_sde_at_warmup=use_sde_at_warmup,
|
use_sde_at_warmup=use_sde_at_warmup,
|
||||||
optimize_memory_usage=optimize_memory_usage,
|
optimize_memory_usage=optimize_memory_usage,
|
||||||
|
supported_action_spaces=(gym.spaces.Box),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.target_entropy = target_entropy
|
self.target_entropy = target_entropy
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
from typing import Any, Dict, List, Optional, Tuple, Type, Union
|
from typing import Any, Dict, List, Optional, Tuple, Type, Union
|
||||||
|
|
||||||
|
import gym
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch as th
|
import torch as th
|
||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
|
|
@ -107,6 +108,7 @@ class TD3(OffPolicyAlgorithm):
|
||||||
seed=seed,
|
seed=seed,
|
||||||
sde_support=False,
|
sde_support=False,
|
||||||
optimize_memory_usage=optimize_memory_usage,
|
optimize_memory_usage=optimize_memory_usage,
|
||||||
|
supported_action_spaces=(gym.spaces.Box),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.policy_delay = policy_delay
|
self.policy_delay = policy_delay
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,7 @@ import gym
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from stable_baselines3 import DQN, SAC, TD3
|
from stable_baselines3 import A2C, DDPG, DQN, PPO, SAC, TD3
|
||||||
from stable_baselines3.common.evaluation import evaluate_policy
|
from stable_baselines3.common.evaluation import evaluate_policy
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -49,3 +49,20 @@ def test_identity_spaces(model_class, env):
|
||||||
model.learn(total_timesteps=500)
|
model.learn(total_timesteps=500)
|
||||||
|
|
||||||
evaluate_policy(model, env, n_eval_episodes=5, warn=False)
|
evaluate_policy(model, env, n_eval_episodes=5, warn=False)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("model_class", [A2C, DDPG, DQN, PPO, SAC, TD3])
|
||||||
|
@pytest.mark.parametrize("env", ["Pendulum-v0", "CartPole-v1"])
|
||||||
|
def test_action_spaces(model_class, env):
|
||||||
|
if model_class in [SAC, DDPG, TD3]:
|
||||||
|
supported_action_space = env == "Pendulum-v0"
|
||||||
|
elif model_class == DQN:
|
||||||
|
supported_action_space = env == "CartPole-v1"
|
||||||
|
elif model_class in [A2C, PPO]:
|
||||||
|
supported_action_space = True
|
||||||
|
|
||||||
|
if supported_action_space:
|
||||||
|
model_class("MlpPolicy", env)
|
||||||
|
else:
|
||||||
|
with pytest.raises(AssertionError):
|
||||||
|
model_class("MlpPolicy", env)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue