diff --git a/alf/config_util.py b/alf/config_util.py index 9addf23bb..5335a88fb 100644 --- a/alf/config_util.py +++ b/alf/config_util.py @@ -15,7 +15,6 @@ from absl import logging import functools -import gin import inspect from inspect import Parameter import os @@ -23,6 +22,10 @@ import runpy import shutil +USE_GIN = os.environ.get('ALF_USE_GIN', "1") == "1" +if USE_GIN: + import gin + __all__ = [ 'config', 'config1', @@ -611,8 +614,7 @@ def _decorate(fn_or_cls, name, whitelist, blacklist): else: fn_or_cls = _make_wrapper(fn_or_cls, configs, signature, has_self=0) - if fn_or_cls.__module__ != '' and os.environ.get( - 'ALF_USE_GIN', "1") == "1": + if fn_or_cls.__module__ != '' and USE_GIN: # If a file is executed using runpy.run_path(), the module name is # '', which is not an acceptable name by gin. return gin.configurable( diff --git a/alf/environments/utils.py b/alf/environments/utils.py index 3f453c1a3..1bdaf03e7 100644 --- a/alf/environments/utils.py +++ b/alf/environments/utils.py @@ -18,7 +18,6 @@ import random import alf -from alf.environments import suite_gym from alf.environments import thread_environment, parallel_environment, fast_parallel_environment from alf.environments import alf_wrappers @@ -89,8 +88,8 @@ def _env_constructor(env_load_fn, env_name, batch_size_per_env, seed, env_id): @alf.configurable -def create_environment(env_name='CartPole-v0', - env_load_fn=suite_gym.load, +def create_environment(env_name=None, + env_load_fn=None, eval_env_load_fn=None, for_evaluation=False, num_parallel_environments=30, @@ -108,10 +107,10 @@ def create_environment(env_name='CartPole-v0', """Create a batched environment. Args: - env_name (str|list[str]): env name. If it is a list, ``MultitaskWrapper`` + env_name (str|list[str]): env name (e.g 'CartPole-v0'). If it is a list, ``MultitaskWrapper`` will be used to create multi-task environments. Each one of them consists of the environments listed in ``env_name``. - env_load_fn (Callable) : callable that create an environment + env_load_fn (Callable) : callable that create an environment (e.g. suite_gym.load) If env_load_fn has attribute ``batched`` and it is True, ``evn_load_fn(env_name, env_id=env_id, batch_size=batch_size_per_env)`` will be used to create the batched environment. Otherwise, @@ -172,6 +171,14 @@ def create_environment(env_name='CartPole-v0', """ + if env_name is None: + # Keep compatibility with the old default env_name + env_name = 'CartPole-v0' + if env_load_fn is None: + # Keep compatibility with the old default env_name + from alf.environments import suite_gym + env_load_fn = suite_gym.load + if for_evaluation: # for creating an evaluation environment, use ``eval_env_load_fn`` if # provided and fall back to ``env_load_fn`` otherwise @@ -266,7 +273,7 @@ def create_environment(env_name='CartPole-v0', @alf.configurable def load_with_random_max_episode_steps(env_name, - env_load_fn=suite_gym.load, + env_load_fn, min_steps=200, max_steps=250): """Create environment with random max_episode_steps in range diff --git a/alf/examples/diayn_pendulum.gin b/alf/examples/diayn_pendulum.gin index 5f735b5ec..b5ba84554 100644 --- a/alf/examples/diayn_pendulum.gin +++ b/alf/examples/diayn_pendulum.gin @@ -4,7 +4,7 @@ import alf.trainers.policy_trainer import alf.algorithms.diayn_algorithm import alf.algorithms.goal_generator import alf.networks.critic_networks - +import alf.environments.suite_gym # skill related num_of_skills=5 diff --git a/alf/examples/hybrid_rl/hybrid_sac_pendulum.gin b/alf/examples/hybrid_rl/hybrid_sac_pendulum.gin index cd2391972..c28045cb5 100644 --- a/alf/examples/hybrid_rl/hybrid_sac_pendulum.gin +++ b/alf/examples/hybrid_rl/hybrid_sac_pendulum.gin @@ -1,4 +1,3 @@ - include 'sac_pendulum.gin' # Example config for Hybrid RL training. diff --git a/alf/examples/icm_mountain_car.gin b/alf/examples/icm_mountain_car.gin index 4ed90738d..93a143ad7 100644 --- a/alf/examples/icm_mountain_car.gin +++ b/alf/examples/icm_mountain_car.gin @@ -2,7 +2,7 @@ import alf.algorithms.agent import alf.algorithms.actor_critic_loss import alf.algorithms.icm_algorithm import alf.algorithms.entropy_target_algorithm - +import alf.environments.suite_gym # environment config create_environment.num_parallel_environments=30 diff --git a/alf/examples/mdq_pendulum.gin b/alf/examples/mdq_pendulum.gin index 083e34766..6e96eb777 100644 --- a/alf/examples/mdq_pendulum.gin +++ b/alf/examples/mdq_pendulum.gin @@ -1,7 +1,7 @@ import alf.utils.math_ops import alf.algorithms.agent import alf.algorithms.mdq_algorithm - +import alf.environments.suite_gym observation_spec=@get_observation_spec() action_spec=@get_action_spec() diff --git a/alf/examples/ppo_bullet_humanoid.gin b/alf/examples/ppo_bullet_humanoid.gin index 9935d09b6..1708f32fb 100644 --- a/alf/examples/ppo_bullet_humanoid.gin +++ b/alf/examples/ppo_bullet_humanoid.gin @@ -1,3 +1,4 @@ +import alf.environments.suite_gym include 'ppo.gin' # environment config diff --git a/alf/examples/sac_pendulum.gin b/alf/examples/sac_pendulum.gin index a9968d3ef..879fd4bf9 100644 --- a/alf/examples/sac_pendulum.gin +++ b/alf/examples/sac_pendulum.gin @@ -1,4 +1,4 @@ - +import alf.environments.suite_gym include 'sac.gin' import alf.utils.math_ops diff --git a/alf/examples/sarsa_ddpg_pendulum.gin b/alf/examples/sarsa_ddpg_pendulum.gin index 3623e7bc4..76659c779 100644 --- a/alf/examples/sarsa_ddpg_pendulum.gin +++ b/alf/examples/sarsa_ddpg_pendulum.gin @@ -1,3 +1,4 @@ +import alf.environments.suite_gym include 'sarsa_ddpg.gin' # environment config diff --git a/alf/examples/sarsa_sac_bipedal_walker.gin b/alf/examples/sarsa_sac_bipedal_walker.gin index 00f41f098..ff1d0d1b6 100644 --- a/alf/examples/sarsa_sac_bipedal_walker.gin +++ b/alf/examples/sarsa_sac_bipedal_walker.gin @@ -1,4 +1,5 @@ include 'sarsa_sac.gin' +import alf.environments.suite_gym # Need to install Box2D by "pip install box2d-py" create_environment.env_name="BipedalWalker-v2" diff --git a/alf/examples/sarsa_sac_pendulum.gin b/alf/examples/sarsa_sac_pendulum.gin index 4a622c0b3..3f1ccb02b 100644 --- a/alf/examples/sarsa_sac_pendulum.gin +++ b/alf/examples/sarsa_sac_pendulum.gin @@ -1,3 +1,4 @@ +import alf.environments.suite_gym include 'sarsa_sac.gin' # environment config