This commit is contained in:
Fabian
2022-07-07 10:47:04 +02:00
parent 4a3134d7be
commit fc00cf8a87
24 changed files with 235 additions and 302 deletions
+6 -8
View File
@@ -20,7 +20,7 @@ def make_dmc(
environment_kwargs: dict = {},
time_limit: Union[None, float] = None,
channels_first: bool = True
):
):
# Adopted from: https://github.com/denisyarats/dmc2gym/blob/master/dmc2gym/__init__.py
# License: MIT
# Copyright (c) 2020 Denis Yarats
@@ -32,12 +32,10 @@ def make_dmc(
env_id = f'dmc_{domain_name}_{task_name}_{seed}-v1'
if from_pixels:
assert not visualize_reward, 'cannot use visualize reward when learning from pixels'
assert not visualize_reward, 'Cannot use visualize reward when learning from pixels.'
# shorten episode length
if episode_length is None:
# Default lengths for benchmarking suite is 1000 and for manipulation tasks 250
episode_length = 250 if domain_name == "manipulation" else 1000
# Default lengths for benchmarking suite is 1000 and for manipulation tasks 250
episode_length = episode_length or (250 if domain_name == "manipulation" else 1000)
max_episode_steps = (episode_length + frame_skip - 1) // frame_skip
if env_id not in gym.envs.registry.env_specs:
@@ -61,7 +59,7 @@ def make_dmc(
camera_id=camera_id,
frame_skip=frame_skip,
channels_first=channels_first,
),
),
max_episode_steps=max_episode_steps,
)
)
return gym.make(env_id)
+27 -15
View File
@@ -8,7 +8,7 @@ from gym.envs.registration import EnvSpec, registry
from gym.wrappers import TimeAwareObservation
from alr_envs.black_box.black_box_wrapper import BlackBoxWrapper
from alr_envs.black_box.controller.controller_factory import get_controller
from alr_envs.black_box.factory.controller_factory import get_controller
from alr_envs.black_box.factory.basis_generator_factory import get_basis_generator
from alr_envs.black_box.factory.phase_generator_factory import get_phase_generator
from alr_envs.black_box.factory.trajectory_generator_factory import get_trajectory_generator
@@ -43,11 +43,7 @@ def make_rank(env_id: str, seed: int, rank: int = 0, return_callable=True, **kwa
def make(env_id, seed, **kwargs):
# This access is required to allow for nested dict updates
spec = registry.get(env_id)
all_kwargs = deepcopy(spec.kwargs)
nested_update(all_kwargs, kwargs)
return _make(env_id, seed, **all_kwargs)
return _make(env_id, seed, **kwargs)
def _make(env_id: str, seed, **kwargs):
@@ -62,12 +58,25 @@ def _make(env_id: str, seed, **kwargs):
Returns: Gym environment
"""
if any(deprec in env_id for deprec in ["DetPMP", "detpmp"]):
warnings.warn("DetPMP is deprecated and converted to ProMP")
env_id = env_id.replace("DetPMP", "ProMP")
env_id = env_id.replace("detpmp", "promp")
# 'dmc:domain-task'
# 'gym:name-vX'
# 'meta:name-vX'
# 'meta:bb:name-vX'
# 'hand:name-vX'
# 'name-vX'
# 'bb:name-vX'
#
# env_id.split(':')
# if 'dmc' :
try:
# This access is required to allow for nested dict updates for BB envs
spec = registry.get(env_id)
all_kwargs = deepcopy(spec.kwargs)
nested_update(all_kwargs, kwargs)
kwargs = all_kwargs
# Add seed to kwargs in case it is a predefined gym+dmc hybrid environment.
if env_id.startswith("dmc"):
kwargs.update({"seed": seed})
@@ -77,22 +86,25 @@ def _make(env_id: str, seed, **kwargs):
env.seed(seed)
env.action_space.seed(seed)
env.observation_space.seed(seed)
except gym.error.Error:
except (gym.error.Error, AttributeError):
# MetaWorld env
import metaworld
if env_id in metaworld.ML1.ENV_NAMES:
env = metaworld.envs.ALL_V2_ENVIRONMENTS_GOAL_OBSERVABLE[env_id + "-goal-observable"](seed=seed, **kwargs)
# setting this avoids generating the same initialization after each reset
env._freeze_rand_vec = False
env.seeded_rand_vec = True
# Manually set spec, as metaworld environments are not registered via gym
env.unwrapped.spec = EnvSpec(env_id)
# Set Timelimit based on the maximum allowed path length of the environment
env = gym.wrappers.TimeLimit(env, max_episode_steps=env.max_path_length)
env.seed(seed)
env.action_space.seed(seed)
env.observation_space.seed(seed)
env.goal_space.seed(seed)
# env.seed(seed)
# env.action_space.seed(seed)
# env.observation_space.seed(seed)
# env.goal_space.seed(seed)
else:
# DMC