unified API wrapper and updated examples

This commit is contained in:
ottofabian
2021-07-02 13:09:56 +02:00
parent 6607d9cff9
commit 80933eba09
22 changed files with 383 additions and 485 deletions
+20 -14
View File
@@ -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(
+24 -10
View File
@@ -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 -3
View File
@@ -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})