unified API wrapper and updated examples
This commit is contained in:
+20
-14
@@ -1,22 +1,24 @@
|
||||
import collections
|
||||
import re
|
||||
from typing import Union
|
||||
|
||||
import gym
|
||||
from gym.envs.registration import register
|
||||
|
||||
|
||||
def make(
|
||||
id,
|
||||
seed=1,
|
||||
visualize_reward=True,
|
||||
from_pixels=False,
|
||||
height=84,
|
||||
width=84,
|
||||
camera_id=0,
|
||||
frame_skip=1,
|
||||
episode_length=1000,
|
||||
environment_kwargs=None,
|
||||
time_limit=None,
|
||||
channels_first=True
|
||||
id: str,
|
||||
seed: int = 1,
|
||||
visualize_reward: bool = True,
|
||||
from_pixels: bool = False,
|
||||
height: int = 84,
|
||||
width: int = 84,
|
||||
camera_id: int = 0,
|
||||
frame_skip: int = 1,
|
||||
episode_length: Union[None, int] = None,
|
||||
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
|
||||
@@ -31,12 +33,16 @@ def make(
|
||||
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
|
||||
|
||||
max_episode_steps = (episode_length + frame_skip - 1) // frame_skip
|
||||
|
||||
if env_id not in gym.envs.registry.env_specs:
|
||||
task_kwargs = {}
|
||||
task_kwargs = {'random': seed}
|
||||
# if seed is not None:
|
||||
task_kwargs['random'] = seed
|
||||
# task_kwargs['random'] = seed
|
||||
if time_limit is not None:
|
||||
task_kwargs['time_limit'] = time_limit
|
||||
register(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# Adopted from: https://github.com/denisyarats/dmc2gym/blob/master/dmc2gym/wrappers.py
|
||||
# License: MIT
|
||||
# Copyright (c) 2020 Denis Yarats
|
||||
import collections
|
||||
from typing import Any, Dict, Tuple
|
||||
|
||||
import numpy as np
|
||||
@@ -31,12 +32,21 @@ def _spec_to_box(spec):
|
||||
return spaces.Box(low, high, dtype=np.float32)
|
||||
|
||||
|
||||
def _flatten_obs(obs):
|
||||
obs_pieces = []
|
||||
for v in obs.values():
|
||||
flat = np.array([v]) if np.isscalar(v) else v.ravel()
|
||||
obs_pieces.append(flat)
|
||||
return np.concatenate(obs_pieces, axis=0)
|
||||
def _flatten_obs(obs: collections.MutableMapping):
|
||||
# obs_pieces = []
|
||||
# for v in obs.values():
|
||||
# flat = np.array([v]) if np.isscalar(v) else v.ravel()
|
||||
# obs_pieces.append(flat)
|
||||
# return np.concatenate(obs_pieces, axis=0)
|
||||
|
||||
if not isinstance(obs, collections.MutableMapping):
|
||||
raise ValueError(f'Requires dict-like observations structure. {type(obs)} found.')
|
||||
|
||||
# Keep key order consistent for non OrderedDicts
|
||||
keys = obs.keys() if isinstance(obs, collections.OrderedDict) else sorted(obs.keys())
|
||||
|
||||
obs_vals = [np.array([obs[key]]) if np.isscalar(obs[key]) else obs[key].ravel() for key in keys]
|
||||
return np.concatenate(obs_vals)
|
||||
|
||||
|
||||
class DMCWrapper(core.Env):
|
||||
@@ -75,7 +85,7 @@ class DMCWrapper(core.Env):
|
||||
self._action_space = _spec_to_box([self._env.action_spec()])
|
||||
self._observation_space = _spec_to_box(self._env.observation_spec().values())
|
||||
|
||||
self._last_observation = None
|
||||
self._last_state = None
|
||||
self.viewer = None
|
||||
|
||||
# set seed
|
||||
@@ -107,6 +117,10 @@ class DMCWrapper(core.Env):
|
||||
def action_space(self):
|
||||
return self._action_space
|
||||
|
||||
@property
|
||||
def dt(self):
|
||||
return self._env.control_timestep() * self._frame_skip
|
||||
|
||||
def seed(self, seed=None):
|
||||
self._action_space.seed(seed)
|
||||
self._observation_space.seed(seed)
|
||||
@@ -123,19 +137,19 @@ class DMCWrapper(core.Env):
|
||||
if done:
|
||||
break
|
||||
|
||||
self._last_observation = _flatten_obs(time_step.observation)
|
||||
self._last_state = _flatten_obs(time_step.observation)
|
||||
obs = self._get_obs(time_step)
|
||||
extra['discount'] = time_step.discount
|
||||
return obs, reward, done, extra
|
||||
|
||||
def reset(self) -> np.ndarray:
|
||||
time_step = self._env.reset()
|
||||
self._last_observation = _flatten_obs(time_step.observation)
|
||||
self._last_state = _flatten_obs(time_step.observation)
|
||||
obs = self._get_obs(time_step)
|
||||
return obs
|
||||
|
||||
def render(self, mode='rgb_array', height=None, width=None, camera_id=0):
|
||||
if self._last_observation is None:
|
||||
if self._last_state is None:
|
||||
raise ValueError('Environment not ready to render. Call reset() first.')
|
||||
|
||||
# assert mode == 'rgb_array', 'only support rgb_array mode, given %s' % mode
|
||||
|
||||
@@ -3,7 +3,7 @@ from typing import Iterable, List, Type
|
||||
|
||||
import gym
|
||||
|
||||
from mp_env_api.env_wrappers.mp_env_wrapper import MPEnvWrapper
|
||||
from mp_env_api.interface_wrappers.mp_env_wrapper import MPEnvWrapper
|
||||
from mp_env_api.mp_wrappers.detpmp_wrapper import DetPMPWrapper
|
||||
from mp_env_api.mp_wrappers.dmp_wrapper import DmpWrapper
|
||||
|
||||
@@ -32,7 +32,7 @@ def make_env_rank(env_id: str, seed: int, rank: int = 0):
|
||||
def make_env(env_id: str, seed, **kwargs):
|
||||
"""
|
||||
Converts an env_id to an environment with the gym API.
|
||||
This also works for DeepMind Control Suite env_wrappers
|
||||
This also works for DeepMind Control Suite interface_wrappers
|
||||
for which domain name and task name are expected to be separated by "-".
|
||||
Args:
|
||||
env_id: gym name or env_id of the form "domain_name-task_name" for DMC tasks
|
||||
@@ -42,7 +42,7 @@ def make_env(env_id: str, seed, **kwargs):
|
||||
|
||||
"""
|
||||
try:
|
||||
# Add seed to kwargs in case it is a predefined dmc environment.
|
||||
# Add seed to kwargs in case it is a predefined gym+dmc hybrid environment.
|
||||
if env_id.startswith("dmc"):
|
||||
kwargs.update({"seed": seed})
|
||||
|
||||
|
||||
Reference in New Issue
Block a user