Fixes for python 2 + env from string

This commit is contained in:
Antonin Raffin 2019-09-06 11:46:25 +02:00
parent 904742714d
commit 90882ee846
3 changed files with 5 additions and 3 deletions

View file

@ -9,7 +9,7 @@ setup(name='torchy_baselines',
install_requires=[ install_requires=[
'gym[classic_control]>=0.10.9', 'gym[classic_control]>=0.10.9',
'numpy', 'numpy',
'torch>=1.2.0' # torch>=1.2.0+cpu 'torch>=1.2.0'
], ],
extras_require={ extras_require={
'tests': [ 'tests': [

View file

@ -5,8 +5,7 @@ import gym
from torchy_baselines import TD3 from torchy_baselines import TD3
def test_pendulum(): def test_pendulum():
env = gym.make("Pendulum-v0") model = TD3('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[64, 64]), start_timesteps=100, verbose=1)
model = TD3('MlpPolicy', env, policy_kwargs=dict(net_arch=[64, 64]), start_timesteps=100, verbose=1)
model.learn(total_timesteps=500, eval_freq=100) model.learn(total_timesteps=500, eval_freq=100)
model.save("test_save") model.save("test_save")
model.load("test_save") model.load("test_save")

View file

@ -34,6 +34,9 @@ class BaseRLModel(object):
self.params = None self.params = None
if env is not None: if env is not None:
if env is not None:
if isinstance(env, str):
env = gym.make(env)
self.env = env self.env = env
self.n_envs = 1 self.n_envs = 1
self.observation_space = env.observation_space self.observation_space = env.observation_space