2020-02-14 13:03:41 +00:00
|
|
|
import gym
|
|
|
|
|
import pytest
|
|
|
|
|
|
2020-06-29 09:16:54 +00:00
|
|
|
from stable_baselines3 import A2C, PPO, SAC, TD3, DQN
|
2020-05-05 13:02:35 +00:00
|
|
|
from stable_baselines3.common.vec_env import DummyVecEnv
|
2020-02-14 13:03:41 +00:00
|
|
|
|
|
|
|
|
MODEL_LIST = [
|
|
|
|
|
PPO,
|
|
|
|
|
A2C,
|
|
|
|
|
TD3,
|
|
|
|
|
SAC,
|
2020-06-29 09:16:54 +00:00
|
|
|
DQN,
|
2020-02-14 13:03:41 +00:00
|
|
|
]
|
|
|
|
|
|
2020-03-12 10:12:10 +00:00
|
|
|
|
2020-02-14 13:03:41 +00:00
|
|
|
@pytest.mark.parametrize("model_class", MODEL_LIST)
|
|
|
|
|
def test_auto_wrap(model_class):
|
|
|
|
|
# test auto wrapping of env into a VecEnv
|
2020-06-29 09:16:54 +00:00
|
|
|
|
|
|
|
|
# Use different environment for DQN
|
|
|
|
|
if model_class is DQN:
|
|
|
|
|
env_name = 'CartPole-v0'
|
|
|
|
|
else:
|
|
|
|
|
env_name = 'Pendulum-v0'
|
|
|
|
|
env = gym.make(env_name)
|
|
|
|
|
eval_env = gym.make(env_name)
|
2020-02-14 13:03:41 +00:00
|
|
|
model = model_class('MlpPolicy', env)
|
|
|
|
|
model.learn(100, eval_env=eval_env)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.parametrize("model_class", MODEL_LIST)
|
2020-02-14 13:15:55 +00:00
|
|
|
@pytest.mark.parametrize("env_id", ['Pendulum-v0', 'CartPole-v1'])
|
|
|
|
|
def test_predict(model_class, env_id):
|
2020-06-29 09:16:54 +00:00
|
|
|
if env_id == 'CartPole-v1':
|
|
|
|
|
if model_class in [SAC, TD3]:
|
|
|
|
|
return
|
|
|
|
|
elif model_class in [DQN]:
|
2020-02-14 13:15:55 +00:00
|
|
|
return
|
2020-02-14 13:03:41 +00:00
|
|
|
|
|
|
|
|
# test detection of different shapes by the predict method
|
2020-02-14 13:15:55 +00:00
|
|
|
model = model_class('MlpPolicy', env_id)
|
|
|
|
|
env = gym.make(env_id)
|
|
|
|
|
vec_env = DummyVecEnv([lambda: gym.make(env_id), lambda: gym.make(env_id)])
|
2020-02-14 13:03:41 +00:00
|
|
|
|
|
|
|
|
obs = env.reset()
|
2020-03-18 14:11:19 +00:00
|
|
|
action, _ = model.predict(obs)
|
2020-02-14 13:15:55 +00:00
|
|
|
assert action.shape == env.action_space.shape
|
2020-02-14 13:03:41 +00:00
|
|
|
assert env.action_space.contains(action)
|
|
|
|
|
|
|
|
|
|
vec_env_obs = vec_env.reset()
|
2020-03-18 14:11:19 +00:00
|
|
|
action, _ = model.predict(vec_env_obs)
|
2020-02-14 13:03:41 +00:00
|
|
|
assert action.shape[0] == vec_env_obs.shape[0]
|