release
This commit is contained in:
Vendored
+229
@@ -0,0 +1,229 @@
|
||||
import os
|
||||
import json
|
||||
|
||||
try:
|
||||
from collections.abc import Iterable
|
||||
except ImportError:
|
||||
Iterable = (tuple, list)
|
||||
|
||||
|
||||
def make_async(
|
||||
id,
|
||||
num_envs=1,
|
||||
asynchronous=True,
|
||||
wrappers=None,
|
||||
render=False,
|
||||
obs_dim=23,
|
||||
action_dim=7,
|
||||
env_type=None,
|
||||
max_episode_steps=None,
|
||||
# below for furniture only
|
||||
gpu_id=0,
|
||||
headless=True,
|
||||
record=False,
|
||||
normalization_path=None,
|
||||
furniture="one_leg",
|
||||
randomness="low",
|
||||
act_steps=8,
|
||||
sparse_reward=False,
|
||||
# below for robomimic only
|
||||
robomimic_env_cfg_path=None,
|
||||
use_image_obs=False,
|
||||
render_offscreen=False,
|
||||
reward_shaping=False,
|
||||
shape_meta=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""Create a vectorized environment from multiple copies of an environment,
|
||||
from its id.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
id : str
|
||||
The environment ID. This must be a valid ID from the registry.
|
||||
|
||||
num_envs : int
|
||||
Number of copies of the environment.
|
||||
|
||||
asynchronous : bool
|
||||
If `True`, wraps the environments in an :class:`AsyncVectorEnv` (which uses
|
||||
`multiprocessing`_ to run the environments in parallel). If ``False``,
|
||||
wraps the environments in a :class:`SyncVectorEnv`.
|
||||
|
||||
wrappers : dictionary, optional
|
||||
Each key is a wrapper class, and each value is a dictionary of arguments
|
||||
|
||||
Returns
|
||||
-------
|
||||
:class:`gym.vector.VectorEnv`
|
||||
The vectorized environment.
|
||||
|
||||
Example
|
||||
-------
|
||||
>>> env = gym.vector.make('CartPole-v1', num_envs=3)
|
||||
>>> env.reset()
|
||||
array([[-0.04456399, 0.04653909, 0.01326909, -0.02099827],
|
||||
[ 0.03073904, 0.00145001, -0.03088818, -0.03131252],
|
||||
[ 0.03468829, 0.01500225, 0.01230312, 0.01825218]],
|
||||
dtype=float32)
|
||||
"""
|
||||
|
||||
if env_type == "furniture":
|
||||
from furniture_bench.envs.observation import DEFAULT_STATE_OBS
|
||||
from furniture_bench.envs.furniture_rl_sim_env import FurnitureRLSimEnv
|
||||
from env.gym_utils.wrapper.furniture import FurnitureRLSimEnvMultiStepWrapper
|
||||
|
||||
env = FurnitureRLSimEnv(
|
||||
act_rot_repr="rot_6d",
|
||||
action_type="pos",
|
||||
april_tags=False,
|
||||
concat_robot_state=True,
|
||||
ctrl_mode="diffik",
|
||||
obs_keys=DEFAULT_STATE_OBS,
|
||||
furniture=furniture,
|
||||
gpu_id=gpu_id,
|
||||
headless=headless,
|
||||
num_envs=num_envs,
|
||||
observation_space="state",
|
||||
randomness=randomness,
|
||||
max_env_steps=max_episode_steps,
|
||||
record=record,
|
||||
pos_scalar=1,
|
||||
rot_scalar=1,
|
||||
stiffness=1_000,
|
||||
damping=200,
|
||||
)
|
||||
|
||||
env = FurnitureRLSimEnvMultiStepWrapper(
|
||||
env,
|
||||
n_obs_steps=1,
|
||||
n_action_steps=act_steps,
|
||||
reward_agg_method="sum",
|
||||
prev_action=False,
|
||||
reset_within_step=False,
|
||||
pass_full_observations=False,
|
||||
normalization_path=normalization_path,
|
||||
sparse_reward=sparse_reward,
|
||||
)
|
||||
|
||||
return env
|
||||
|
||||
# avoid import error due incompatible gym versions
|
||||
from gym import spaces
|
||||
from env.gym_utils.async_vector_env import AsyncVectorEnv
|
||||
from env.gym_utils.sync_vector_env import SyncVectorEnv
|
||||
from env.gym_utils.wrapper import wrapper_dict
|
||||
|
||||
__all__ = [
|
||||
"AsyncVectorEnv",
|
||||
"SyncVectorEnv",
|
||||
"VectorEnv",
|
||||
"VectorEnvWrapper",
|
||||
"make",
|
||||
]
|
||||
|
||||
# import the envs
|
||||
if robomimic_env_cfg_path is not None:
|
||||
import robomimic.utils.env_utils as EnvUtils
|
||||
import robomimic.utils.obs_utils as ObsUtils
|
||||
elif "avoiding" in id:
|
||||
import gym_avoiding
|
||||
else:
|
||||
import d4rl.gym_mujoco
|
||||
from gym.envs import make as make_
|
||||
|
||||
def _make_env():
|
||||
if robomimic_env_cfg_path is not None:
|
||||
obs_modality_dict = {
|
||||
"low_dim": (
|
||||
wrappers.robomimic_image.low_dim_keys
|
||||
if "robomimic_image" in wrappers
|
||||
else wrappers.robomimic_lowdim.low_dim_keys
|
||||
),
|
||||
"rgb": (
|
||||
wrappers.robomimic_image.image_keys
|
||||
if "robomimic_image" in wrappers
|
||||
else None
|
||||
),
|
||||
}
|
||||
if obs_modality_dict["rgb"] is None:
|
||||
obs_modality_dict.pop("rgb")
|
||||
ObsUtils.initialize_obs_modality_mapping_from_dict(obs_modality_dict)
|
||||
if render_offscreen or use_image_obs:
|
||||
os.environ["MUJOCO_GL"] = "egl"
|
||||
with open(robomimic_env_cfg_path, "r") as f:
|
||||
env_meta = json.load(f)
|
||||
env_meta["reward_shaping"] = reward_shaping
|
||||
env = EnvUtils.create_env_from_metadata(
|
||||
env_meta=env_meta,
|
||||
render=render,
|
||||
# only way to not show collision geometry is to enable render_offscreen, which uses a lot of RAM.
|
||||
render_offscreen=render_offscreen,
|
||||
use_image_obs=use_image_obs,
|
||||
# render_gpu_device_id=0,
|
||||
)
|
||||
# Robosuite's hard reset causes excessive memory consumption.
|
||||
# Disabled to run more envs.
|
||||
# https://github.com/ARISE-Initiative/robosuite/blob/92abf5595eddb3a845cd1093703e5a3ccd01e77e/robosuite/environments/base.py#L247-L248
|
||||
env.env.hard_reset = False
|
||||
else: # d3il, gym
|
||||
env = make_(id, render=render, **kwargs)
|
||||
|
||||
# add wrappers
|
||||
if wrappers is not None:
|
||||
for wrapper, args in wrappers.items():
|
||||
env = wrapper_dict[wrapper](env, **args)
|
||||
return env
|
||||
|
||||
def dummy_env_fn():
|
||||
"""TODO(allenzren): does this dummy env allow camera obs for other envs besides robomimic?"""
|
||||
import gym
|
||||
import numpy as np
|
||||
from env.gym_utils.wrapper.multi_step import MultiStep
|
||||
|
||||
# Avoid importing or using env in the main process
|
||||
# to prevent OpenGL context issue with fork.
|
||||
# Create a fake env whose sole purpose is to provide
|
||||
# obs/action spaces and metadata.
|
||||
env = gym.Env()
|
||||
if shape_meta is not None: # rn only for images
|
||||
observation_space = spaces.Dict()
|
||||
for key, value in shape_meta["obs"].items():
|
||||
shape = value["shape"]
|
||||
if key.endswith("rgb"):
|
||||
min_value, max_value = -1, 1
|
||||
elif key.endswith("state"):
|
||||
min_value, max_value = -1, 1
|
||||
else:
|
||||
raise RuntimeError(f"Unsupported type {key}")
|
||||
this_space = spaces.Box(
|
||||
low=min_value,
|
||||
high=max_value,
|
||||
shape=shape,
|
||||
dtype=np.float32,
|
||||
)
|
||||
observation_space[key] = this_space
|
||||
env.observation_space = observation_space
|
||||
else:
|
||||
env.observation_space = gym.spaces.Box(
|
||||
-1, 1, shape=(obs_dim,), dtype=np.float64
|
||||
)
|
||||
env.action_space = gym.spaces.Box(-1, 1, shape=(action_dim,), dtype=np.int64)
|
||||
env.metadata = {
|
||||
"render.modes": ["human", "rgb_array", "depth_array"],
|
||||
"video.frames_per_second": 12,
|
||||
}
|
||||
return MultiStep(env=env) # use all defaults
|
||||
|
||||
env_fns = [_make_env for _ in range(num_envs)]
|
||||
return (
|
||||
AsyncVectorEnv(
|
||||
env_fns,
|
||||
dummy_env_fn=(
|
||||
dummy_env_fn if render or render_offscreen or use_image_obs else None
|
||||
),
|
||||
delay_init="avoiding" in id, # add delay for D3IL initialization
|
||||
)
|
||||
if asynchronous
|
||||
else SyncVectorEnv(env_fns)
|
||||
)
|
||||
Vendored
+839
@@ -0,0 +1,839 @@
|
||||
"""
|
||||
From gym==0.22.0
|
||||
|
||||
Disable auto-reset after done.
|
||||
Add reset_arg() that allows all environments with different options.
|
||||
Add reset_one_arg() that allows resetting a single environment with options.
|
||||
Add render().
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional, Union, List
|
||||
|
||||
import numpy as np
|
||||
import multiprocessing as mp
|
||||
import time
|
||||
import sys
|
||||
from enum import Enum
|
||||
from copy import deepcopy
|
||||
import time
|
||||
|
||||
from gym import logger
|
||||
|
||||
# from gym.vector.vector_env import VectorEnv
|
||||
from .vector_env import VectorEnv
|
||||
from gym.error import (
|
||||
AlreadyPendingCallError,
|
||||
NoAsyncCallError,
|
||||
ClosedEnvironmentError,
|
||||
CustomSpaceError,
|
||||
)
|
||||
from gym.vector.utils import (
|
||||
create_shared_memory,
|
||||
create_empty_array,
|
||||
write_to_shared_memory,
|
||||
read_from_shared_memory,
|
||||
concatenate,
|
||||
iterate,
|
||||
CloudpickleWrapper,
|
||||
clear_mpi_env_vars,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["AsyncVectorEnv"]
|
||||
|
||||
|
||||
class AsyncState(Enum):
|
||||
DEFAULT = "default"
|
||||
WAITING_RESET = "reset"
|
||||
WAITING_STEP = "step"
|
||||
WAITING_CALL = "call"
|
||||
|
||||
|
||||
class AsyncVectorEnv(VectorEnv):
|
||||
"""Vectorized environment that runs multiple environments in parallel. It
|
||||
uses `multiprocessing`_ processes, and pipes for communication.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
env_fns : iterable of callable
|
||||
Functions that create the environments.
|
||||
|
||||
observation_space : :class:`gym.spaces.Space`, optional
|
||||
Observation space of a single environment. If ``None``, then the
|
||||
observation space of the first environment is taken.
|
||||
|
||||
action_space : :class:`gym.spaces.Space`, optional
|
||||
Action space of a single environment. If ``None``, then the action space
|
||||
of the first environment is taken.
|
||||
|
||||
shared_memory : bool
|
||||
If ``True``, then the observations from the worker processes are
|
||||
communicated back through shared variables. This can improve the
|
||||
efficiency if the observations are large (e.g. images).
|
||||
|
||||
copy : bool
|
||||
If ``True``, then the :meth:`~AsyncVectorEnv.reset` and
|
||||
:meth:`~AsyncVectorEnv.step` methods return a copy of the observations.
|
||||
|
||||
context : str, optional
|
||||
Context for `multiprocessing`_. If ``None``, then the default context is used.
|
||||
|
||||
daemon : bool
|
||||
If ``True``, then subprocesses have ``daemon`` flag turned on; that is, they
|
||||
will quit if the head process quits. However, ``daemon=True`` prevents
|
||||
subprocesses to spawn children, so for some environments you may want
|
||||
to have it set to ``False``.
|
||||
|
||||
worker : callable, optional
|
||||
If set, then use that worker in a subprocess instead of a default one.
|
||||
Can be useful to override some inner vector env logic, for instance,
|
||||
how resets on done are handled.
|
||||
|
||||
Warning
|
||||
-------
|
||||
:attr:`worker` is an advanced mode option. It provides a high degree of
|
||||
flexibility and a high chance to shoot yourself in the foot; thus,
|
||||
if you are writing your own worker, it is recommended to start from the code
|
||||
for ``_worker`` (or ``_worker_shared_memory``) method, and add changes.
|
||||
|
||||
Raises
|
||||
------
|
||||
RuntimeError
|
||||
If the observation space of some sub-environment does not match
|
||||
:obj:`observation_space` (or, by default, the observation space of
|
||||
the first sub-environment).
|
||||
|
||||
ValueError
|
||||
If :obj:`observation_space` is a custom space (i.e. not a default
|
||||
space in Gym, such as :class:`~gym.spaces.Box`, :class:`~gym.spaces.Discrete`,
|
||||
or :class:`~gym.spaces.Dict`) and :obj:`shared_memory` is ``True``.
|
||||
|
||||
Example
|
||||
-------
|
||||
|
||||
.. code-block::
|
||||
|
||||
>>> env = gym.vector.AsyncVectorEnv([
|
||||
... lambda: gym.make("Pendulum-v0", g=9.81),
|
||||
... lambda: gym.make("Pendulum-v0", g=1.62)
|
||||
... ])
|
||||
>>> env.reset()
|
||||
array([[-0.8286432 , 0.5597771 , 0.90249056],
|
||||
[-0.85009176, 0.5266346 , 0.60007906]], dtype=float32)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
env_fns,
|
||||
dummy_env_fn=None,
|
||||
observation_space=None,
|
||||
action_space=None,
|
||||
shared_memory=True,
|
||||
copy=True,
|
||||
context=None,
|
||||
daemon=True,
|
||||
worker=None,
|
||||
delay_init=False,
|
||||
):
|
||||
ctx = mp.get_context(context)
|
||||
self.env_fns = env_fns
|
||||
self.shared_memory = shared_memory
|
||||
self.copy = copy
|
||||
if dummy_env_fn is None:
|
||||
dummy_env_fn = env_fns[0]
|
||||
dummy_env = dummy_env_fn()
|
||||
self.metadata = dummy_env.metadata
|
||||
|
||||
self.n_envs = len(env_fns)
|
||||
if (observation_space is None) or (action_space is None):
|
||||
observation_space = observation_space or dummy_env.observation_space
|
||||
action_space = action_space or dummy_env.action_space
|
||||
dummy_env.close()
|
||||
del dummy_env
|
||||
super().__init__(
|
||||
num_envs=len(env_fns),
|
||||
observation_space=observation_space,
|
||||
action_space=action_space,
|
||||
)
|
||||
|
||||
if self.shared_memory:
|
||||
try:
|
||||
_obs_buffer = create_shared_memory(
|
||||
self.single_observation_space, n=self.num_envs, ctx=ctx
|
||||
)
|
||||
self.observations = read_from_shared_memory(
|
||||
self.single_observation_space, _obs_buffer, n=self.num_envs
|
||||
)
|
||||
except CustomSpaceError:
|
||||
raise ValueError(
|
||||
"Using `shared_memory=True` in `AsyncVectorEnv` "
|
||||
"is incompatible with non-standard Gym observation spaces "
|
||||
"(i.e. custom spaces inheriting from `gym.Space`), and is "
|
||||
"only compatible with default Gym spaces (e.g. `Box`, "
|
||||
"`Tuple`, `Dict`) for batching. Set `shared_memory=False` "
|
||||
"if you use custom observation spaces."
|
||||
)
|
||||
else:
|
||||
_obs_buffer = None
|
||||
self.observations = create_empty_array(
|
||||
self.single_observation_space, n=self.num_envs, fn=np.zeros
|
||||
)
|
||||
|
||||
self.parent_pipes, self.processes = [], []
|
||||
self.error_queue = ctx.Queue()
|
||||
target = _worker_shared_memory if self.shared_memory else _worker
|
||||
target = worker or target
|
||||
with clear_mpi_env_vars():
|
||||
for idx, env_fn in enumerate(self.env_fns):
|
||||
parent_pipe, child_pipe = ctx.Pipe()
|
||||
process = ctx.Process(
|
||||
target=target,
|
||||
name=f"Worker<{type(self).__name__}>-{idx}",
|
||||
args=(
|
||||
idx,
|
||||
CloudpickleWrapper(env_fn),
|
||||
child_pipe,
|
||||
parent_pipe,
|
||||
_obs_buffer,
|
||||
self.error_queue,
|
||||
),
|
||||
)
|
||||
|
||||
self.parent_pipes.append(parent_pipe)
|
||||
self.processes.append(process)
|
||||
|
||||
process.daemon = daemon
|
||||
process.start()
|
||||
child_pipe.close()
|
||||
if (
|
||||
delay_init
|
||||
): # D3IL complains about temporary XML if n_envs is too large. Adding a delay avoids the error.
|
||||
time.sleep(0.1)
|
||||
|
||||
self._state = AsyncState.DEFAULT
|
||||
# self._check_spaces()
|
||||
|
||||
def seed(self, seed=None):
|
||||
super().seed(seed=seed)
|
||||
self._assert_is_running()
|
||||
if seed is None:
|
||||
seed = [None for _ in range(self.num_envs)]
|
||||
if isinstance(seed, int):
|
||||
seed = [seed + i for i in range(self.num_envs)]
|
||||
assert len(seed) == self.num_envs
|
||||
|
||||
if self._state != AsyncState.DEFAULT:
|
||||
raise AlreadyPendingCallError(
|
||||
f"Calling `seed` while waiting for a pending call to `{self._state.value}` to complete.",
|
||||
self._state.value,
|
||||
)
|
||||
|
||||
for pipe, seed in zip(self.parent_pipes, seed):
|
||||
pipe.send(("seed", seed))
|
||||
_, successes = zip(*[pipe.recv() for pipe in self.parent_pipes])
|
||||
self._raise_if_errors(successes)
|
||||
|
||||
def reset_async(
|
||||
self,
|
||||
seed: Optional[Union[int, List[int]]] = None,
|
||||
return_info: bool = False,
|
||||
options: Optional[dict] = None,
|
||||
):
|
||||
"""Send the calls to :obj:`reset` to each sub-environment.
|
||||
|
||||
Raises
|
||||
------
|
||||
ClosedEnvironmentError
|
||||
If the environment was closed (if :meth:`close` was previously called).
|
||||
|
||||
AlreadyPendingCallError
|
||||
If the environment is already waiting for a pending call to another
|
||||
method (e.g. :meth:`step_async`). This can be caused by two consecutive
|
||||
calls to :meth:`reset_async`, with no call to :meth:`reset_wait` in
|
||||
between.
|
||||
"""
|
||||
self._assert_is_running()
|
||||
|
||||
if seed is None:
|
||||
seed = [None for _ in range(self.num_envs)]
|
||||
if isinstance(seed, int):
|
||||
seed = [seed + i for i in range(self.num_envs)]
|
||||
assert len(seed) == self.num_envs
|
||||
|
||||
if self._state != AsyncState.DEFAULT:
|
||||
raise AlreadyPendingCallError(
|
||||
f"Calling `reset_async` while waiting for a pending call to `{self._state.value}` to complete",
|
||||
self._state.value,
|
||||
)
|
||||
|
||||
for pipe, single_seed in zip(self.parent_pipes, seed):
|
||||
single_kwargs = {}
|
||||
if single_seed is not None:
|
||||
single_kwargs["seed"] = single_seed
|
||||
if return_info:
|
||||
single_kwargs["return_info"] = return_info
|
||||
if options is not None:
|
||||
single_kwargs["options"] = options
|
||||
|
||||
pipe.send(("reset", single_kwargs))
|
||||
self._state = AsyncState.WAITING_RESET
|
||||
|
||||
def reset_wait(
|
||||
self,
|
||||
timeout=None,
|
||||
seed: Optional[int] = None,
|
||||
return_info: bool = False,
|
||||
options: Optional[dict] = None,
|
||||
):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
timeout : int or float, optional
|
||||
Number of seconds before the call to `reset_wait` times out. If
|
||||
`None`, the call to `reset_wait` never times out.
|
||||
seed: ignored
|
||||
options: ignored
|
||||
|
||||
Returns
|
||||
-------
|
||||
element of :attr:`~VectorEnv.observation_space`
|
||||
A batch of observations from the vectorized environment.
|
||||
infos : list of dicts containing metadata
|
||||
|
||||
Raises
|
||||
------
|
||||
ClosedEnvironmentError
|
||||
If the environment was closed (if :meth:`close` was previously called).
|
||||
|
||||
NoAsyncCallError
|
||||
If :meth:`reset_wait` was called without any prior call to
|
||||
:meth:`reset_async`.
|
||||
|
||||
TimeoutError
|
||||
If :meth:`reset_wait` timed out.
|
||||
"""
|
||||
self._assert_is_running()
|
||||
if self._state != AsyncState.WAITING_RESET:
|
||||
raise NoAsyncCallError(
|
||||
"Calling `reset_wait` without any prior " "call to `reset_async`.",
|
||||
AsyncState.WAITING_RESET.value,
|
||||
)
|
||||
|
||||
if not self._poll(timeout):
|
||||
self._state = AsyncState.DEFAULT
|
||||
raise mp.TimeoutError(
|
||||
f"The call to `reset_wait` has timed out after {timeout} second(s)."
|
||||
)
|
||||
|
||||
results, successes = zip(*[pipe.recv() for pipe in self.parent_pipes])
|
||||
self._raise_if_errors(successes)
|
||||
self._state = AsyncState.DEFAULT
|
||||
|
||||
if return_info:
|
||||
results, infos = zip(*results)
|
||||
infos = list(infos)
|
||||
|
||||
if not self.shared_memory:
|
||||
self.observations = concatenate(
|
||||
self.single_observation_space, results, self.observations
|
||||
)
|
||||
|
||||
return (
|
||||
deepcopy(self.observations) if self.copy else self.observations
|
||||
), infos
|
||||
else:
|
||||
if not self.shared_memory:
|
||||
self.observations = concatenate(
|
||||
self.single_observation_space, results, self.observations
|
||||
)
|
||||
|
||||
return deepcopy(self.observations) if self.copy else self.observations
|
||||
|
||||
def step_async(self, actions):
|
||||
"""Send the calls to :obj:`step` to each sub-environment.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
actions : element of :attr:`~VectorEnv.action_space`
|
||||
Batch of actions.
|
||||
|
||||
Raises
|
||||
------
|
||||
ClosedEnvironmentError
|
||||
If the environment was closed (if :meth:`close` was previously called).
|
||||
|
||||
AlreadyPendingCallError
|
||||
If the environment is already waiting for a pending call to another
|
||||
method (e.g. :meth:`reset_async`). This can be caused by two consecutive
|
||||
calls to :meth:`step_async`, with no call to :meth:`step_wait` in
|
||||
between.
|
||||
"""
|
||||
self._assert_is_running()
|
||||
if self._state != AsyncState.DEFAULT:
|
||||
raise AlreadyPendingCallError(
|
||||
f"Calling `step_async` while waiting for a pending call to `{self._state.value}` to complete.",
|
||||
self._state.value,
|
||||
)
|
||||
|
||||
actions = iterate(self.action_space, actions)
|
||||
for pipe, action in zip(self.parent_pipes, actions):
|
||||
pipe.send(("step", action))
|
||||
self._state = AsyncState.WAITING_STEP
|
||||
|
||||
def step_wait(self, timeout=None):
|
||||
"""Wait for the calls to :obj:`step` in each sub-environment to finish.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
timeout : int or float, optional
|
||||
Number of seconds before the call to :meth:`step_wait` times out. If
|
||||
``None``, the call to :meth:`step_wait` never times out.
|
||||
|
||||
Returns
|
||||
-------
|
||||
observations : element of :attr:`~VectorEnv.observation_space`
|
||||
A batch of observations from the vectorized environment.
|
||||
|
||||
rewards : :obj:`np.ndarray`, dtype :obj:`np.float_`
|
||||
A vector of rewards from the vectorized environment.
|
||||
|
||||
dones : :obj:`np.ndarray`, dtype :obj:`np.bool_`
|
||||
A vector whose entries indicate whether the episode has ended.
|
||||
|
||||
infos : list of dict
|
||||
A list of auxiliary diagnostic information dicts from sub-environments.
|
||||
|
||||
Raises
|
||||
------
|
||||
ClosedEnvironmentError
|
||||
If the environment was closed (if :meth:`close` was previously called).
|
||||
|
||||
NoAsyncCallError
|
||||
If :meth:`step_wait` was called without any prior call to
|
||||
:meth:`step_async`.
|
||||
|
||||
TimeoutError
|
||||
If :meth:`step_wait` timed out.
|
||||
"""
|
||||
self._assert_is_running()
|
||||
if self._state != AsyncState.WAITING_STEP:
|
||||
raise NoAsyncCallError(
|
||||
"Calling `step_wait` without any prior call " "to `step_async`.",
|
||||
AsyncState.WAITING_STEP.value,
|
||||
)
|
||||
|
||||
if not self._poll(timeout):
|
||||
self._state = AsyncState.DEFAULT
|
||||
raise mp.TimeoutError(
|
||||
f"The call to `step_wait` has timed out after {timeout} second(s)."
|
||||
)
|
||||
|
||||
results, successes = zip(*[pipe.recv() for pipe in self.parent_pipes])
|
||||
self._raise_if_errors(successes)
|
||||
self._state = AsyncState.DEFAULT
|
||||
observations_list, rewards, dones, infos = zip(*results)
|
||||
|
||||
if not self.shared_memory:
|
||||
self.observations = concatenate(
|
||||
self.single_observation_space,
|
||||
observations_list,
|
||||
self.observations,
|
||||
)
|
||||
|
||||
return (
|
||||
deepcopy(self.observations) if self.copy else self.observations,
|
||||
np.array(rewards),
|
||||
np.array(dones, dtype=np.bool_),
|
||||
infos,
|
||||
)
|
||||
|
||||
def call_async(self, name, *args, **kwargs):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
name : string
|
||||
Name of the method or property to call.
|
||||
|
||||
*args
|
||||
Arguments to apply to the method call.
|
||||
|
||||
**kwargs
|
||||
Keywoard arguments to apply to the method call.
|
||||
"""
|
||||
self._assert_is_running()
|
||||
if self._state != AsyncState.DEFAULT:
|
||||
raise AlreadyPendingCallError(
|
||||
"Calling `call_async` while waiting "
|
||||
f"for a pending call to `{self._state.value}` to complete.",
|
||||
self._state.value,
|
||||
)
|
||||
|
||||
for pipe in self.parent_pipes:
|
||||
pipe.send(("_call", (name, args, kwargs)))
|
||||
self._state = AsyncState.WAITING_CALL
|
||||
|
||||
def call_wait(self, timeout=None):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
timeout : int or float, optional
|
||||
Number of seconds before the call to `step_wait` times out. If
|
||||
`None` (default), the call to `step_wait` never times out.
|
||||
|
||||
Returns
|
||||
-------
|
||||
results : list
|
||||
List of the results of the individual calls to the method or
|
||||
property for each environment.
|
||||
"""
|
||||
self._assert_is_running()
|
||||
if self._state != AsyncState.WAITING_CALL:
|
||||
raise NoAsyncCallError(
|
||||
"Calling `call_wait` without any prior call to `call_async`.",
|
||||
AsyncState.WAITING_CALL.value,
|
||||
)
|
||||
|
||||
if not self._poll(timeout):
|
||||
self._state = AsyncState.DEFAULT
|
||||
raise mp.TimeoutError(
|
||||
f"The call to `call_wait` has timed out after {timeout} second(s)."
|
||||
)
|
||||
|
||||
results, successes = zip(*[pipe.recv() for pipe in self.parent_pipes])
|
||||
self._raise_if_errors(successes)
|
||||
self._state = AsyncState.DEFAULT
|
||||
|
||||
return results
|
||||
|
||||
def set_attr(self, name, values):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
name : string
|
||||
Name of the property to be set in each individual environment.
|
||||
|
||||
values : list, tuple, or object
|
||||
Values of the property to be set to. If `values` is a list or
|
||||
tuple, then it corresponds to the values for each individual
|
||||
environment, otherwise a single value is set for all environments.
|
||||
"""
|
||||
self._assert_is_running()
|
||||
if not isinstance(values, (list, tuple)):
|
||||
values = [values for _ in range(self.num_envs)]
|
||||
if len(values) != self.num_envs:
|
||||
raise ValueError(
|
||||
"Values must be a list or tuple with length equal to the "
|
||||
f"number of environments. Got `{len(values)}` values for "
|
||||
f"{self.num_envs} environments."
|
||||
)
|
||||
|
||||
if self._state != AsyncState.DEFAULT:
|
||||
raise AlreadyPendingCallError(
|
||||
"Calling `set_attr` while waiting "
|
||||
f"for a pending call to `{self._state.value}` to complete.",
|
||||
self._state.value,
|
||||
)
|
||||
|
||||
for pipe, value in zip(self.parent_pipes, values):
|
||||
pipe.send(("_setattr", (name, value)))
|
||||
_, successes = zip(*[pipe.recv() for pipe in self.parent_pipes])
|
||||
self._raise_if_errors(successes)
|
||||
|
||||
def close_extras(self, timeout=None, terminate=False):
|
||||
"""Close the environments & clean up the extra resources
|
||||
(processes and pipes).
|
||||
|
||||
Parameters
|
||||
----------
|
||||
timeout : int or float, optional
|
||||
Number of seconds before the call to :meth:`close` times out. If ``None``,
|
||||
the call to :meth:`close` never times out. If the call to :meth:`close`
|
||||
times out, then all processes are terminated.
|
||||
|
||||
terminate : bool
|
||||
If ``True``, then the :meth:`close` operation is forced and all processes
|
||||
are terminated.
|
||||
|
||||
Raises
|
||||
------
|
||||
TimeoutError
|
||||
If :meth:`close` timed out.
|
||||
"""
|
||||
timeout = 0 if terminate else timeout
|
||||
try:
|
||||
if self._state != AsyncState.DEFAULT:
|
||||
logger.warn(
|
||||
f"Calling `close` while waiting for a pending call to `{self._state.value}` to complete."
|
||||
)
|
||||
function = getattr(self, f"{self._state.value}_wait")
|
||||
function(timeout)
|
||||
except mp.TimeoutError:
|
||||
terminate = True
|
||||
|
||||
if terminate:
|
||||
for process in self.processes:
|
||||
if process.is_alive():
|
||||
process.terminate()
|
||||
else:
|
||||
for pipe in self.parent_pipes:
|
||||
if (pipe is not None) and (not pipe.closed):
|
||||
pipe.send(("close", None))
|
||||
for pipe in self.parent_pipes:
|
||||
if (pipe is not None) and (not pipe.closed):
|
||||
pipe.recv()
|
||||
|
||||
for pipe in self.parent_pipes:
|
||||
if pipe is not None:
|
||||
pipe.close()
|
||||
for process in self.processes:
|
||||
process.join()
|
||||
|
||||
def _poll(self, timeout=None):
|
||||
self._assert_is_running()
|
||||
if timeout is None:
|
||||
return True
|
||||
end_time = time.perf_counter() + timeout
|
||||
delta = None
|
||||
for pipe in self.parent_pipes:
|
||||
delta = max(end_time - time.perf_counter(), 0)
|
||||
if pipe is None:
|
||||
return False
|
||||
if pipe.closed or (not pipe.poll(delta)):
|
||||
return False
|
||||
return True
|
||||
|
||||
def _check_spaces(self):
|
||||
self._assert_is_running()
|
||||
spaces = (self.single_observation_space, self.single_action_space)
|
||||
for pipe in self.parent_pipes:
|
||||
pipe.send(("_check_spaces", spaces))
|
||||
results, successes = zip(*[pipe.recv() for pipe in self.parent_pipes])
|
||||
self._raise_if_errors(successes)
|
||||
same_observation_spaces, same_action_spaces = zip(*results)
|
||||
if not all(same_observation_spaces):
|
||||
raise RuntimeError(
|
||||
"Some environments have an observation space different from "
|
||||
f"`{self.single_observation_space}`. In order to batch observations, "
|
||||
"the observation spaces from all environments must be equal."
|
||||
)
|
||||
if not all(same_action_spaces):
|
||||
raise RuntimeError(
|
||||
"Some environments have an action space different from "
|
||||
f"`{self.single_action_space}`. In order to batch actions, the "
|
||||
"action spaces from all environments must be equal."
|
||||
)
|
||||
|
||||
def _assert_is_running(self):
|
||||
if self.closed:
|
||||
raise ClosedEnvironmentError(
|
||||
f"Trying to operate on `{type(self).__name__}`, after a call to `close()`."
|
||||
)
|
||||
|
||||
def _raise_if_errors(self, successes):
|
||||
if all(successes):
|
||||
return
|
||||
|
||||
num_errors = self.num_envs - sum(successes)
|
||||
assert num_errors > 0
|
||||
for _ in range(num_errors):
|
||||
index, exctype, value = self.error_queue.get()
|
||||
logger.error(
|
||||
f"Received the following error from Worker-{index}: {exctype.__name__}: {value}"
|
||||
)
|
||||
logger.error(f"Shutting down Worker-{index}.")
|
||||
self.parent_pipes[index].close()
|
||||
self.parent_pipes[index] = None
|
||||
|
||||
logger.error("Raising the last exception back to the main process.")
|
||||
raise exctype(value)
|
||||
|
||||
def __del__(self):
|
||||
if not getattr(self, "closed", True):
|
||||
self.close(terminate=True)
|
||||
|
||||
######################### Added #########################
|
||||
def call_sync(self, method_name, indices=None, **method_kwargs):
|
||||
"""Call instance methods of vectorized environments."""
|
||||
target_remotes = self._get_target_remotes(indices)
|
||||
for remote in target_remotes:
|
||||
remote.send(("_call_sync", (method_name, method_kwargs)))
|
||||
return [remote.recv() for remote in target_remotes]
|
||||
|
||||
def call_sync_arg(
|
||||
self, method_name, method_arg_name, method_arg_list, indices=None
|
||||
):
|
||||
"""Call instance methods of vectorized environments with args."""
|
||||
target_remotes = self._get_target_remotes(indices)
|
||||
for method_arg, remote in zip(method_arg_list, target_remotes):
|
||||
method_kwargs = {method_arg_name: method_arg}
|
||||
remote.send(("_call_sync", (method_name, method_kwargs)))
|
||||
return [remote.recv() for remote in target_remotes]
|
||||
|
||||
def _get_target_remotes(self, indices):
|
||||
"""Get the connection object needed to communicate with the wanted
|
||||
envs that are in subprocesses."""
|
||||
if indices is None:
|
||||
indices = range(self.n_envs)
|
||||
return [self.parent_pipes[i] for i in indices]
|
||||
|
||||
def reset_arg(self, options_list, **kwargs):
|
||||
results = self.call_sync_arg("reset", "options", options_list)
|
||||
obs = [result[0] for result in results]
|
||||
if isinstance(obs[0], np.ndarray):
|
||||
return np.stack(obs)
|
||||
else:
|
||||
assert isinstance(obs[0], dict)
|
||||
return obs
|
||||
|
||||
def reset_one_arg(self, env_ind, options=None):
|
||||
"""
|
||||
Reset one environment with options.
|
||||
"""
|
||||
obs, success = self.call_sync(
|
||||
"reset",
|
||||
options=options,
|
||||
indices=[env_ind],
|
||||
)[0]
|
||||
return obs
|
||||
|
||||
def render(self, *args, **kwargs):
|
||||
return self.call("render", *args, **kwargs)
|
||||
|
||||
|
||||
def _worker(index, env_fn, pipe, parent_pipe, shared_memory, error_queue):
|
||||
assert shared_memory is None
|
||||
env = env_fn()
|
||||
parent_pipe.close()
|
||||
try:
|
||||
while True:
|
||||
command, data = pipe.recv()
|
||||
if command == "reset":
|
||||
if "return_info" in data and data["return_info"] == True:
|
||||
observation, info = env.reset(**data)
|
||||
pipe.send(((observation, info), True))
|
||||
else:
|
||||
observation = env.reset(**data)
|
||||
pipe.send((observation, True))
|
||||
|
||||
elif command == "step":
|
||||
observation, reward, done, info = env.step(data)
|
||||
# if done:
|
||||
# info["terminal_observation"] = observation
|
||||
# observation = env.reset()
|
||||
pipe.send(((observation, reward, done, info), True))
|
||||
elif command == "seed":
|
||||
env.seed(data)
|
||||
pipe.send((None, True))
|
||||
elif command == "close":
|
||||
pipe.send((None, True))
|
||||
break
|
||||
elif command == "_call_sync":
|
||||
function = getattr(env, data[0])
|
||||
pipe.send((function(**data[1]), True))
|
||||
elif command == "_call":
|
||||
name, args, kwargs = data
|
||||
if name in ["reset", "step", "seed", "close"]:
|
||||
raise ValueError(
|
||||
f"Trying to call function `{name}` with "
|
||||
f"`_call`. Use `{name}` directly instead."
|
||||
)
|
||||
function = getattr(env, name)
|
||||
if callable(function):
|
||||
pipe.send((function(**kwargs), True))
|
||||
else:
|
||||
pipe.send((function, True))
|
||||
elif command == "_setattr":
|
||||
name, value = data
|
||||
setattr(env, name, value)
|
||||
pipe.send((None, True))
|
||||
elif command == "_check_spaces":
|
||||
pipe.send(
|
||||
(
|
||||
(data[0] == env.observation_space, data[1] == env.action_space),
|
||||
True,
|
||||
)
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Received unknown command `{command}`. Must "
|
||||
"be one of {`reset`, `step`, `seed`, `close`, `_call`, "
|
||||
"`_setattr`, `_check_spaces`}."
|
||||
)
|
||||
except (KeyboardInterrupt, Exception):
|
||||
error_queue.put((index,) + sys.exc_info()[:2])
|
||||
pipe.send((None, False))
|
||||
finally:
|
||||
env.close()
|
||||
|
||||
|
||||
def _worker_shared_memory(index, env_fn, pipe, parent_pipe, shared_memory, error_queue):
|
||||
assert shared_memory is not None
|
||||
env = env_fn()
|
||||
observation_space = env.observation_space
|
||||
parent_pipe.close()
|
||||
try:
|
||||
while True:
|
||||
command, data = pipe.recv()
|
||||
if command == "reset":
|
||||
if "return_info" in data and data["return_info"] == True:
|
||||
observation, info = env.reset(**data)
|
||||
write_to_shared_memory(
|
||||
observation_space, index, observation, shared_memory
|
||||
)
|
||||
pipe.send(((None, info), True))
|
||||
else:
|
||||
observation = env.reset(**data)
|
||||
write_to_shared_memory(
|
||||
observation_space, index, observation, shared_memory
|
||||
)
|
||||
pipe.send((None, True))
|
||||
elif command == "step":
|
||||
observation, reward, done, info = env.step(data)
|
||||
# if done:
|
||||
# info["terminal_observation"] = observation
|
||||
# observation = env.reset()
|
||||
write_to_shared_memory(
|
||||
observation_space, index, observation, shared_memory
|
||||
)
|
||||
pipe.send(((None, reward, done, info), True))
|
||||
elif command == "seed":
|
||||
env.seed(data)
|
||||
pipe.send((None, True))
|
||||
elif command == "close":
|
||||
pipe.send((None, True))
|
||||
break
|
||||
elif command == "_call_sync":
|
||||
function = getattr(env, data[0])
|
||||
pipe.send((function(**data[1]), True))
|
||||
elif command == "_call":
|
||||
name, args, kwargs = data
|
||||
if name in ["reset", "step", "seed", "close"]:
|
||||
raise ValueError(
|
||||
f"Trying to call function `{name}` with "
|
||||
f"`_call`. Use `{name}` directly instead."
|
||||
)
|
||||
function = getattr(env, name)
|
||||
if callable(function):
|
||||
pipe.send((function(*args, **kwargs), True))
|
||||
else:
|
||||
pipe.send((function, True))
|
||||
elif command == "_setattr":
|
||||
name, value = data
|
||||
setattr(env, name, value)
|
||||
pipe.send((None, True))
|
||||
elif command == "_check_spaces":
|
||||
pipe.send(
|
||||
((data[0] == observation_space, data[1] == env.action_space), True)
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Received unknown command `{command}`. Must "
|
||||
"be one of {`reset`, `step`, `seed`, `close`, `_call`, "
|
||||
"`_setattr`, `_check_spaces`}."
|
||||
)
|
||||
except (KeyboardInterrupt, Exception):
|
||||
error_queue.put((index,) + sys.exc_info()[:2])
|
||||
pipe.send((None, False))
|
||||
finally:
|
||||
env.close()
|
||||
+80
@@ -0,0 +1,80 @@
|
||||
"""
|
||||
Normalization for Furniture-Bench environments.
|
||||
|
||||
TODO: use this normalizer for all benchmarks.
|
||||
|
||||
"""
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class LinearNormalizer(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.stats = nn.ParameterDict()
|
||||
|
||||
def fit(self, data_dict):
|
||||
for key, tensor in data_dict.items():
|
||||
min_value = tensor.min(dim=0)[0]
|
||||
max_value = tensor.max(dim=0)[0]
|
||||
|
||||
# Check if any column has only one value throughout
|
||||
diff = max_value - min_value
|
||||
constant_columns = diff == 0
|
||||
|
||||
# Set a small range for constant columns to avoid division by zero
|
||||
min_value[constant_columns] -= 1
|
||||
max_value[constant_columns] += 1
|
||||
|
||||
self.stats[key] = nn.ParameterDict(
|
||||
{
|
||||
"min": nn.Parameter(min_value, requires_grad=False),
|
||||
"max": nn.Parameter(max_value, requires_grad=False),
|
||||
},
|
||||
)
|
||||
self._turn_off_gradients()
|
||||
|
||||
def _normalize(self, x, key):
|
||||
stats = self.stats[key]
|
||||
x = (x - stats["min"]) / (stats["max"] - stats["min"])
|
||||
x = 2 * x - 1
|
||||
return x
|
||||
|
||||
def _denormalize(self, x, key):
|
||||
stats = self.stats[key]
|
||||
x = (x + 1) / 2
|
||||
x = x * (stats["max"] - stats["min"]) + stats["min"]
|
||||
return x
|
||||
|
||||
def forward(self, x, key, forward=True):
|
||||
if forward:
|
||||
return self._normalize(x, key)
|
||||
else:
|
||||
return self._denormalize(x, key)
|
||||
|
||||
def _turn_off_gradients(self):
|
||||
for key in self.stats.keys():
|
||||
for stat in self.stats[key].keys():
|
||||
self.stats[key][stat].requires_grad = False
|
||||
|
||||
def load_state_dict(self, state_dict):
|
||||
|
||||
stats = nn.ParameterDict()
|
||||
for key, value in state_dict.items():
|
||||
if key.startswith("stats."):
|
||||
param_key = key[6:]
|
||||
keys = param_key.split(".")
|
||||
current_dict = stats
|
||||
for k in keys[:-1]:
|
||||
if k not in current_dict:
|
||||
current_dict[k] = nn.ParameterDict()
|
||||
current_dict = current_dict[k]
|
||||
current_dict[keys[-1]] = nn.Parameter(value)
|
||||
|
||||
self.stats = stats
|
||||
self._turn_off_gradients()
|
||||
|
||||
return f"<Added keys {self.stats.keys()} to the normalizer.>"
|
||||
|
||||
def keys(self):
|
||||
return self.stats.keys()
|
||||
Vendored
+201
@@ -0,0 +1,201 @@
|
||||
from typing import List, Union, Optional
|
||||
|
||||
import numpy as np
|
||||
from copy import deepcopy
|
||||
|
||||
from gym import logger
|
||||
from gym.logger import warn
|
||||
from gym.vector.vector_env import VectorEnv
|
||||
from gym.vector.utils import concatenate, iterate, create_empty_array
|
||||
|
||||
|
||||
__all__ = ["SyncVectorEnv"]
|
||||
|
||||
|
||||
class SyncVectorEnv(VectorEnv):
|
||||
"""Vectorized environment that serially runs multiple environments.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
env_fns : iterable of callable
|
||||
Functions that create the environments.
|
||||
|
||||
observation_space : :class:`gym.spaces.Space`, optional
|
||||
Observation space of a single environment. If ``None``, then the
|
||||
observation space of the first environment is taken.
|
||||
|
||||
action_space : :class:`gym.spaces.Space`, optional
|
||||
Action space of a single environment. If ``None``, then the action space
|
||||
of the first environment is taken.
|
||||
|
||||
copy : bool
|
||||
If ``True``, then the :meth:`reset` and :meth:`step` methods return a
|
||||
copy of the observations.
|
||||
|
||||
Raises
|
||||
------
|
||||
RuntimeError
|
||||
If the observation space of some sub-environment does not match
|
||||
:obj:`observation_space` (or, by default, the observation space of
|
||||
the first sub-environment).
|
||||
|
||||
Example
|
||||
-------
|
||||
|
||||
.. code-block::
|
||||
|
||||
>>> env = gym.vector.SyncVectorEnv([
|
||||
... lambda: gym.make("Pendulum-v0", g=9.81),
|
||||
... lambda: gym.make("Pendulum-v0", g=1.62)
|
||||
... ])
|
||||
>>> env.reset()
|
||||
array([[-0.8286432 , 0.5597771 , 0.90249056],
|
||||
[-0.85009176, 0.5266346 , 0.60007906]], dtype=float32)
|
||||
"""
|
||||
|
||||
def __init__(self, env_fns, observation_space=None, action_space=None, copy=True):
|
||||
self.env_fns = env_fns
|
||||
self.envs = [env_fn() for env_fn in env_fns]
|
||||
self.copy = copy
|
||||
self.metadata = self.envs[0].metadata
|
||||
|
||||
if (observation_space is None) or (action_space is None):
|
||||
observation_space = observation_space or self.envs[0].observation_space
|
||||
action_space = action_space or self.envs[0].action_space
|
||||
super().__init__(
|
||||
num_envs=len(env_fns),
|
||||
observation_space=observation_space,
|
||||
action_space=action_space,
|
||||
)
|
||||
|
||||
self._check_spaces()
|
||||
self.observations = create_empty_array(
|
||||
self.single_observation_space, n=self.num_envs, fn=np.zeros
|
||||
)
|
||||
self._rewards = np.zeros((self.num_envs,), dtype=np.float64)
|
||||
self._dones = np.zeros((self.num_envs,), dtype=np.bool_)
|
||||
self._actions = None
|
||||
|
||||
def seed(self, seed=None):
|
||||
super().seed(seed=seed)
|
||||
if seed is None:
|
||||
seed = [None for _ in range(self.num_envs)]
|
||||
if isinstance(seed, int):
|
||||
seed = [seed + i for i in range(self.num_envs)]
|
||||
assert len(seed) == self.num_envs
|
||||
|
||||
for env, single_seed in zip(self.envs, seed):
|
||||
env.seed(single_seed)
|
||||
|
||||
def reset_wait(
|
||||
self,
|
||||
seed: Optional[Union[int, List[int]]] = None,
|
||||
return_info: bool = False,
|
||||
options: Optional[dict] = None,
|
||||
):
|
||||
if seed is None:
|
||||
seed = [None for _ in range(self.num_envs)]
|
||||
if isinstance(seed, int):
|
||||
seed = [seed + i for i in range(self.num_envs)]
|
||||
assert len(seed) == self.num_envs
|
||||
|
||||
self._dones[:] = False
|
||||
observations = []
|
||||
data_list = []
|
||||
for env, single_seed in zip(self.envs, seed):
|
||||
|
||||
kwargs = {}
|
||||
if single_seed is not None:
|
||||
kwargs["seed"] = single_seed
|
||||
if options is not None:
|
||||
kwargs["options"] = options
|
||||
if return_info == True:
|
||||
kwargs["return_info"] = return_info
|
||||
|
||||
if not return_info:
|
||||
observation = env.reset(**kwargs)
|
||||
observations.append(observation)
|
||||
else:
|
||||
observation, data = env.reset(**kwargs)
|
||||
observations.append(observation)
|
||||
data_list.append(data)
|
||||
|
||||
self.observations = concatenate(
|
||||
self.single_observation_space, observations, self.observations
|
||||
)
|
||||
if not return_info:
|
||||
return deepcopy(self.observations) if self.copy else self.observations
|
||||
else:
|
||||
return (
|
||||
deepcopy(self.observations) if self.copy else self.observations
|
||||
), data_list
|
||||
|
||||
def step_async(self, actions):
|
||||
self._actions = iterate(self.action_space, actions)
|
||||
|
||||
def step_wait(self):
|
||||
observations, infos = [], []
|
||||
for i, (env, action) in enumerate(zip(self.envs, self._actions)):
|
||||
observation, self._rewards[i], self._dones[i], info = env.step(action)
|
||||
if self._dones[i]:
|
||||
info["terminal_observation"] = observation
|
||||
observation = env.reset()
|
||||
observations.append(observation)
|
||||
infos.append(info)
|
||||
self.observations = concatenate(
|
||||
self.single_observation_space, observations, self.observations
|
||||
)
|
||||
|
||||
return (
|
||||
deepcopy(self.observations) if self.copy else self.observations,
|
||||
np.copy(self._rewards),
|
||||
np.copy(self._dones),
|
||||
infos,
|
||||
)
|
||||
|
||||
def call(self, name, *args, **kwargs):
|
||||
results = []
|
||||
for env in self.envs:
|
||||
function = getattr(env, name)
|
||||
if callable(function):
|
||||
results.append(function(*args, **kwargs))
|
||||
else:
|
||||
results.append(function)
|
||||
|
||||
return tuple(results)
|
||||
|
||||
def set_attr(self, name, values):
|
||||
if not isinstance(values, (list, tuple)):
|
||||
values = [values for _ in range(self.num_envs)]
|
||||
if len(values) != self.num_envs:
|
||||
raise ValueError(
|
||||
"Values must be a list or tuple with length equal to the "
|
||||
f"number of environments. Got `{len(values)}` values for "
|
||||
f"{self.num_envs} environments."
|
||||
)
|
||||
|
||||
for env, value in zip(self.envs, values):
|
||||
setattr(env, name, value)
|
||||
|
||||
def close_extras(self, **kwargs):
|
||||
"""Close the environments."""
|
||||
[env.close() for env in self.envs]
|
||||
|
||||
def _check_spaces(self):
|
||||
for env in self.envs:
|
||||
if not (env.observation_space == self.single_observation_space):
|
||||
raise RuntimeError(
|
||||
"Some environments have an observation space different from "
|
||||
f"`{self.single_observation_space}`. In order to batch observations, "
|
||||
"the observation spaces from all environments must be equal."
|
||||
)
|
||||
|
||||
if not (env.action_space == self.single_action_space):
|
||||
raise RuntimeError(
|
||||
"Some environments have an action space different from "
|
||||
f"`{self.single_action_space}`. In order to batch actions, the "
|
||||
"action spaces from all environments must be equal."
|
||||
)
|
||||
|
||||
else:
|
||||
return True
|
||||
Vendored
+276
@@ -0,0 +1,276 @@
|
||||
from typing import Optional, Union, List
|
||||
|
||||
import gym
|
||||
from gym.logger import warn, deprecation
|
||||
from gym.spaces import Tuple
|
||||
from gym.vector.utils.spaces import batch_space
|
||||
|
||||
|
||||
__all__ = ["VectorEnv"]
|
||||
|
||||
|
||||
class VectorEnv(gym.Env):
|
||||
r"""Base class for vectorized environments.
|
||||
|
||||
Each observation returned from vectorized environment is a batch of observations
|
||||
for each sub-environment. And :meth:`step` is also expected to receive a batch of
|
||||
actions for each sub-environment.
|
||||
|
||||
.. note::
|
||||
|
||||
All sub-environments should share the identical observation and action spaces.
|
||||
In other words, a vector of multiple different environments is not supported.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
num_envs : int
|
||||
Number of environments in the vectorized environment.
|
||||
|
||||
observation_space : :class:`gym.spaces.Space`
|
||||
Observation space of a single environment.
|
||||
|
||||
action_space : :class:`gym.spaces.Space`
|
||||
Action space of a single environment.
|
||||
"""
|
||||
|
||||
def __init__(self, num_envs, observation_space, action_space):
|
||||
self.num_envs = num_envs
|
||||
self.is_vector_env = True
|
||||
self.observation_space = batch_space(observation_space, n=num_envs)
|
||||
self.action_space = batch_space(action_space, n=num_envs)
|
||||
|
||||
self.closed = False
|
||||
self.viewer = None
|
||||
|
||||
# The observation and action spaces of a single environment are
|
||||
# kept in separate properties
|
||||
self.single_observation_space = observation_space
|
||||
self.single_action_space = action_space
|
||||
|
||||
def reset_async(
|
||||
self,
|
||||
seed: Optional[Union[int, List[int]]] = None,
|
||||
return_info: bool = False,
|
||||
options: Optional[dict] = None,
|
||||
):
|
||||
pass
|
||||
|
||||
def reset_wait(
|
||||
self,
|
||||
seed: Optional[Union[int, List[int]]] = None,
|
||||
return_info: bool = False,
|
||||
options: Optional[dict] = None,
|
||||
):
|
||||
raise NotImplementedError()
|
||||
|
||||
def reset(
|
||||
self,
|
||||
*,
|
||||
seed: Optional[Union[int, List[int]]] = None,
|
||||
return_info: bool = False,
|
||||
options: Optional[dict] = None,
|
||||
):
|
||||
r"""Reset all sub-environments and return a batch of initial observations.
|
||||
|
||||
Returns
|
||||
-------
|
||||
element of :attr:`observation_space`
|
||||
A batch of observations from the vectorized environment.
|
||||
"""
|
||||
self.reset_async(seed=seed, return_info=return_info, options=options)
|
||||
return self.reset_wait(seed=seed, return_info=return_info, options=options)
|
||||
|
||||
def step_async(self, actions):
|
||||
pass
|
||||
|
||||
def step_wait(self, **kwargs):
|
||||
raise NotImplementedError()
|
||||
|
||||
def step(self, actions):
|
||||
r"""Take an action for each sub-environments.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
actions : element of :attr:`action_space`
|
||||
Batch of actions.
|
||||
|
||||
Returns
|
||||
-------
|
||||
observations : element of :attr:`observation_space`
|
||||
A batch of observations from the vectorized environment.
|
||||
|
||||
rewards : :obj:`np.ndarray`, dtype :obj:`np.float_`
|
||||
A vector of rewards from the vectorized environment.
|
||||
|
||||
dones : :obj:`np.ndarray`, dtype :obj:`np.bool_`
|
||||
A vector whose entries indicate whether the episode has ended.
|
||||
|
||||
infos : list of dict
|
||||
A list of auxiliary diagnostic information dicts from sub-environments.
|
||||
"""
|
||||
|
||||
self.step_async(actions)
|
||||
return self.step_wait()
|
||||
|
||||
def call_async(self, name, *args, **kwargs):
|
||||
pass
|
||||
|
||||
def call_wait(self, **kwargs):
|
||||
raise NotImplementedError()
|
||||
|
||||
def call(self, name, *args, **kwargs):
|
||||
"""Call a method, or get a property, from each sub-environment.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name : string
|
||||
Name of the method or property to call.
|
||||
|
||||
*args
|
||||
Arguments to apply to the method call.
|
||||
|
||||
**kwargs
|
||||
Keywoard arguments to apply to the method call.
|
||||
|
||||
Returns
|
||||
-------
|
||||
results : list
|
||||
List of the results of the individual calls to the method or
|
||||
property for each environment.
|
||||
"""
|
||||
self.call_async(name, *args, **kwargs)
|
||||
return self.call_wait()
|
||||
|
||||
def get_attr(self, name):
|
||||
"""Get a property from each sub-environment.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name : string
|
||||
Name of the property to be get from each individual environment.
|
||||
"""
|
||||
return self.call(name)
|
||||
|
||||
def set_attr(self, name, values):
|
||||
"""Set a property in each sub-environment.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name : string
|
||||
Name of the property to be set in each individual environment.
|
||||
|
||||
values : list, tuple, or object
|
||||
Values of the property to be set to. If `values` is a list or
|
||||
tuple, then it corresponds to the values for each individual
|
||||
environment, otherwise a single value is set for all environments.
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
def close_extras(self, **kwargs):
|
||||
r"""Clean up the extra resources e.g. beyond what's in this base class."""
|
||||
pass
|
||||
|
||||
def close(self, **kwargs):
|
||||
r"""Close all sub-environments and release resources.
|
||||
|
||||
It also closes all the existing image viewers, then calls :meth:`close_extras` and set
|
||||
:attr:`closed` as ``True``.
|
||||
|
||||
.. warning::
|
||||
|
||||
This function itself does not close the environments, it should be handled
|
||||
in :meth:`close_extras`. This is generic for both synchronous and asynchronous
|
||||
vectorized environments.
|
||||
|
||||
.. note::
|
||||
|
||||
This will be automatically called when garbage collected or program exited.
|
||||
|
||||
"""
|
||||
if self.closed:
|
||||
return
|
||||
if self.viewer is not None:
|
||||
self.viewer.close()
|
||||
self.close_extras(**kwargs)
|
||||
self.closed = True
|
||||
|
||||
def seed(self, seed=None):
|
||||
"""Set the random seed in all sub-environments.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
seed : list of int, or int, optional
|
||||
Random seed for each sub-environment. If ``seed`` is a list of
|
||||
length ``num_envs``, then the items of the list are chosen as random
|
||||
seeds. If ``seed`` is an int, then each sub-environment uses the random
|
||||
seed ``seed + n``, where ``n`` is the index of the sub-environment
|
||||
(between ``0`` and ``num_envs - 1``).
|
||||
"""
|
||||
deprecation(
|
||||
"Function `env.seed(seed)` is marked as deprecated and will be removed in the future. "
|
||||
"Please use `env.reset(seed=seed) instead in VectorEnvs."
|
||||
)
|
||||
|
||||
def __del__(self):
|
||||
if not getattr(self, "closed", True):
|
||||
self.close()
|
||||
|
||||
def __repr__(self):
|
||||
if self.spec is None:
|
||||
return f"{self.__class__.__name__}({self.num_envs})"
|
||||
else:
|
||||
return f"{self.__class__.__name__}({self.spec.id}, {self.num_envs})"
|
||||
|
||||
|
||||
class VectorEnvWrapper(VectorEnv):
|
||||
r"""Wraps the vectorized environment to allow a modular transformation.
|
||||
|
||||
This class is the base class for all wrappers for vectorized environments. The subclass
|
||||
could override some methods to change the behavior of the original vectorized environment
|
||||
without touching the original code.
|
||||
|
||||
.. note::
|
||||
|
||||
Don't forget to call ``super().__init__(env)`` if the subclass overrides :meth:`__init__`.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, env):
|
||||
assert isinstance(env, VectorEnv)
|
||||
self.env = env
|
||||
|
||||
# explicitly forward the methods defined in VectorEnv
|
||||
# to self.env (instead of the base class)
|
||||
def reset_async(self, **kwargs):
|
||||
return self.env.reset_async(**kwargs)
|
||||
|
||||
def reset_wait(self, **kwargs):
|
||||
return self.env.reset_wait(**kwargs)
|
||||
|
||||
def step_async(self, actions):
|
||||
return self.env.step_async(actions)
|
||||
|
||||
def step_wait(self):
|
||||
return self.env.step_wait()
|
||||
|
||||
def close(self, **kwargs):
|
||||
return self.env.close(**kwargs)
|
||||
|
||||
def close_extras(self, **kwargs):
|
||||
return self.env.close_extras(**kwargs)
|
||||
|
||||
def seed(self, seed=None):
|
||||
return self.env.seed(seed)
|
||||
|
||||
# implicitly forward all other methods and attributes to self.env
|
||||
def __getattr__(self, name):
|
||||
if name.startswith("_"):
|
||||
raise AttributeError(f"attempted to get missing private attribute '{name}'")
|
||||
return getattr(self.env, name)
|
||||
|
||||
@property
|
||||
def unwrapped(self):
|
||||
return self.env.unwrapped
|
||||
|
||||
def __repr__(self):
|
||||
return f"<{self.__class__.__name__}, {self.env}>"
|
||||
Vendored
+14
@@ -0,0 +1,14 @@
|
||||
from .multi_step import MultiStep
|
||||
from .robomimic_lowdim import RobomimicLowdimWrapper
|
||||
from .robomimic_image import RobomimicImageWrapper
|
||||
from .d3il_lowdim import D3ilLowdimWrapper
|
||||
from .mujoco_locomotion_lowdim import MujocoLocomotionLowdimWrapper
|
||||
|
||||
|
||||
wrapper_dict = {
|
||||
"multi_step": MultiStep,
|
||||
"robomimic_lowdim": RobomimicLowdimWrapper,
|
||||
"robomimic_image": RobomimicImageWrapper,
|
||||
"d3il_lowdim": D3ilLowdimWrapper,
|
||||
"mujoco_locomotion_lowdim": MujocoLocomotionLowdimWrapper,
|
||||
}
|
||||
Vendored
+87
@@ -0,0 +1,87 @@
|
||||
"""
|
||||
Environment wrapper for D3IL environments with state observations.
|
||||
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import gym
|
||||
|
||||
|
||||
class D3ilLowdimWrapper(gym.Env):
|
||||
def __init__(
|
||||
self,
|
||||
env,
|
||||
normalization_path,
|
||||
# init_state=None,
|
||||
# render_hw=(256, 256),
|
||||
# render_camera_name="agentview",
|
||||
):
|
||||
self.env = env
|
||||
# self.init_state = init_state
|
||||
# self.render_hw = render_hw
|
||||
# self.render_camera_name = render_camera_name
|
||||
|
||||
# setup spaces
|
||||
self.action_space = env.action_space
|
||||
self.observation_space = env.observation_space
|
||||
normalization = np.load(normalization_path)
|
||||
self.obs_min = normalization["obs_min"]
|
||||
self.obs_max = normalization["obs_max"]
|
||||
self.action_min = normalization["action_min"]
|
||||
self.action_max = normalization["action_max"]
|
||||
|
||||
# def get_observation(self):
|
||||
# raw_obs = self.env.get_observation()
|
||||
# obs = np.concatenate([raw_obs[key] for key in self.obs_keys], axis=0)
|
||||
# return obs
|
||||
|
||||
def seed(self, seed=None):
|
||||
if seed is not None:
|
||||
np.random.seed(seed=seed)
|
||||
else:
|
||||
np.random.seed()
|
||||
|
||||
def reset(self, **kwargs):
|
||||
"""Ignore passed-in arguments like seed"""
|
||||
options = kwargs.get("options", {})
|
||||
|
||||
new_seed = options.get(
|
||||
"seed", None
|
||||
) # used to set all environments to specified seeds
|
||||
# if self.init_state is not None:
|
||||
# # always reset to the same state to be compatible with gym
|
||||
# self.env.reset_to({"states": self.init_state})
|
||||
if new_seed is not None:
|
||||
self.seed(seed=new_seed)
|
||||
obs = self.env.reset()
|
||||
else:
|
||||
# random reset
|
||||
obs = self.env.reset()
|
||||
|
||||
# normalize
|
||||
obs = self.normalize_obs(obs)
|
||||
return obs
|
||||
|
||||
def normalize_obs(self, obs):
|
||||
return 2 * ((obs - self.obs_min) / (self.obs_max - self.obs_min + 1e-6) - 0.5)
|
||||
|
||||
def unnormaliza_action(self, action):
|
||||
action = (action + 1) / 2 # [-1, 1] -> [0, 1]
|
||||
return action * (self.action_max - self.action_min) + self.action_min
|
||||
|
||||
def step(self, action):
|
||||
action = self.unnormaliza_action(action)
|
||||
obs, reward, done, info = self.env.step(action)
|
||||
|
||||
# normalize
|
||||
obs = self.normalize_obs(obs)
|
||||
return obs, reward, done, info
|
||||
|
||||
def render(self, mode="rgb_array"):
|
||||
h, w = self.render_hw
|
||||
return self.env.render(
|
||||
mode=mode,
|
||||
height=h,
|
||||
width=w,
|
||||
camera_name=self.render_camera_name,
|
||||
)
|
||||
Vendored
+152
@@ -0,0 +1,152 @@
|
||||
"""
|
||||
Environment wrapper for Furniture-Bench environments.
|
||||
|
||||
"""
|
||||
|
||||
import gym
|
||||
import numpy as np
|
||||
from furniture_bench.envs.furniture_rl_sim_env import FurnitureRLSimEnv
|
||||
import torch
|
||||
from furniture_bench.controllers.control_utils import proprioceptive_quat_to_6d_rotation
|
||||
from ..furniture_normalizer import LinearNormalizer
|
||||
from .multi_step import repeated_space
|
||||
|
||||
import logging
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class FurnitureRLSimEnvMultiStepWrapper(gym.Wrapper):
|
||||
env: FurnitureRLSimEnv
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
env: FurnitureRLSimEnv,
|
||||
n_obs_steps=1,
|
||||
n_action_steps=1,
|
||||
max_episode_steps=None,
|
||||
sparse_reward=False,
|
||||
reward_agg_method="sum", # never use other types
|
||||
reset_within_step=False,
|
||||
pass_full_observations=False,
|
||||
normalization_path=None,
|
||||
prev_action=False,
|
||||
):
|
||||
assert (
|
||||
not reset_within_step
|
||||
), "reset_within_step must be False for furniture envs"
|
||||
assert n_obs_steps == 1, "n_obs_steps must be 1"
|
||||
assert reward_agg_method == "sum", "reward_agg_method must be sum"
|
||||
assert (
|
||||
not pass_full_observations
|
||||
), "pass_full_observations is not implemented yet"
|
||||
assert not prev_action, "prev_action is not implemented yet"
|
||||
|
||||
super().__init__(env)
|
||||
self._single_action_space = env.action_space
|
||||
self._action_space = repeated_space(env.action_space, n_action_steps)
|
||||
self._observation_space = repeated_space(env.observation_space, n_obs_steps)
|
||||
self.max_episode_steps = max_episode_steps
|
||||
self.n_obs_steps = n_obs_steps
|
||||
self.n_action_steps = n_action_steps
|
||||
self.pass_full_observations = pass_full_observations
|
||||
|
||||
# Use the original reward function where the robot does not receive new reward after completing one part
|
||||
self.sparse_reward = sparse_reward
|
||||
|
||||
# set up normalization
|
||||
self.normalize = normalization_path is not None
|
||||
self.normalizer = LinearNormalizer()
|
||||
self.normalizer.load_state_dict(
|
||||
torch.load(normalization_path, map_location=self.device, weights_only=True)
|
||||
)
|
||||
log.info(f"Loaded normalization from {normalization_path}")
|
||||
|
||||
def reset(
|
||||
self,
|
||||
**kwargs,
|
||||
):
|
||||
"""Resets the environment."""
|
||||
obs = self.env.reset()
|
||||
nobs = self.process_obs(obs)
|
||||
self.best_reward = torch.zeros(self.env.num_envs).to(self.device)
|
||||
self.done = list()
|
||||
|
||||
return nobs
|
||||
|
||||
def reset_arg(self, options_list=None):
|
||||
return self.reset()
|
||||
|
||||
def reset_one_arg(self, env_ind=None, options=None):
|
||||
if env_ind is not None:
|
||||
env_ind = torch.tensor([env_ind], device=self.device)
|
||||
|
||||
return self.reset()
|
||||
|
||||
def step(self, action: np.ndarray):
|
||||
"""
|
||||
Takes in a chunk of actions of length n_action_steps
|
||||
and steps the environment n_action_steps times
|
||||
and returns an aggregated observation, reward, and done signal
|
||||
"""
|
||||
# action: (n_envs, n_action_steps, action_dim)
|
||||
action = torch.tensor(action, device=self.device)
|
||||
|
||||
# Denormalize the action
|
||||
action = self.normalizer(action, "actions", forward=False)
|
||||
|
||||
# Step the environment n_action_steps times
|
||||
obs, sparse_reward, dense_reward, done, info = self._inner_step(action)
|
||||
if self.sparse_reward:
|
||||
reward = sparse_reward.clone().cpu().numpy()
|
||||
else:
|
||||
reward = dense_reward.clone().cpu().numpy()
|
||||
|
||||
# Only mark the environment as done if it times out, ignore done from inner steps
|
||||
truncated = self.env.env_steps >= self.max_env_steps
|
||||
done = truncated
|
||||
|
||||
nobs: np.ndarray = self.process_obs(obs)
|
||||
done: np.ndarray = done.squeeze().cpu().numpy()
|
||||
|
||||
return (nobs, reward, done, info)
|
||||
|
||||
def _inner_step(self, action_chunk: torch.Tensor):
|
||||
dones = torch.zeros(
|
||||
action_chunk.shape[0], dtype=torch.bool, device=action_chunk.device
|
||||
)
|
||||
dense_reward = torch.zeros(action_chunk.shape[0], device=action_chunk.device)
|
||||
sparse_reward = torch.zeros(action_chunk.shape[0], device=action_chunk.device)
|
||||
for i in range(self.n_action_steps):
|
||||
# The dimensions of the action_chunk are (num_envs, chunk_size, action_dim)
|
||||
obs, reward, done, info = self.env.step(action_chunk[:, i, :])
|
||||
|
||||
# track raw reward
|
||||
sparse_reward += reward.squeeze()
|
||||
|
||||
# track best reward --- reward nonzero only one part is assembled
|
||||
self.best_reward += reward.squeeze()
|
||||
|
||||
# assign "permanent" rewards
|
||||
dense_reward += self.best_reward
|
||||
|
||||
dones = dones | done.squeeze()
|
||||
|
||||
return obs, sparse_reward, dense_reward, dones, info
|
||||
|
||||
def process_obs(self, obs: torch.Tensor) -> np.ndarray:
|
||||
robot_state = obs["robot_state"]
|
||||
|
||||
# Convert the robot state to have 6D pose
|
||||
robot_state = proprioceptive_quat_to_6d_rotation(robot_state)
|
||||
|
||||
parts_poses = obs["parts_poses"]
|
||||
|
||||
obs = torch.cat([robot_state, parts_poses], dim=-1)
|
||||
nobs = self.normalizer(obs, "observations", forward=True)
|
||||
nobs = torch.clamp(nobs, -5, 5)
|
||||
|
||||
# Insert a dummy dimension for the n_obs_steps (n_envs, obs_dim) -> (n_envs, n_obs_steps, obs_dim)
|
||||
nobs = nobs.unsqueeze(1).cpu().numpy()
|
||||
|
||||
return nobs
|
||||
@@ -0,0 +1,61 @@
|
||||
"""
|
||||
Environment wrapper for Gym environments (MuJoCo locomotion tasks) with state observations.
|
||||
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import gym
|
||||
|
||||
|
||||
class MujocoLocomotionLowdimWrapper(gym.Env):
|
||||
def __init__(
|
||||
self,
|
||||
env,
|
||||
normalization_path,
|
||||
):
|
||||
self.env = env
|
||||
|
||||
# setup spaces
|
||||
self.action_space = env.action_space
|
||||
self.observation_space = env.observation_space
|
||||
normalization = np.load(normalization_path)
|
||||
self.obs_min = normalization["obs_min"]
|
||||
self.obs_max = normalization["obs_max"]
|
||||
self.action_min = normalization["action_min"]
|
||||
self.action_max = normalization["action_max"]
|
||||
|
||||
def seed(self, seed=None):
|
||||
if seed is not None:
|
||||
np.random.seed(seed=seed)
|
||||
else:
|
||||
np.random.seed()
|
||||
|
||||
def reset(self, **kwargs):
|
||||
"""Ignore passed-in arguments like seed"""
|
||||
options = kwargs.get("options", {})
|
||||
new_seed = options.get("seed", None)
|
||||
if new_seed is not None:
|
||||
self.seed(seed=new_seed)
|
||||
raw_obs = self.env.reset()
|
||||
|
||||
# normalize
|
||||
obs = self.normalize_obs(raw_obs)
|
||||
return obs
|
||||
|
||||
def normalize_obs(self, obs):
|
||||
return 2 * ((obs - self.obs_min) / (self.obs_max - self.obs_min + 1e-6) - 0.5)
|
||||
|
||||
def unnormaliza_action(self, action):
|
||||
action = (action + 1) / 2 # [-1, 1] -> [0, 1]
|
||||
return action * (self.action_max - self.action_min) + self.action_min
|
||||
|
||||
def step(self, action):
|
||||
raw_action = self.unnormaliza_action(action)
|
||||
raw_obs, reward, done, info = self.env.step(raw_action)
|
||||
|
||||
# normalize
|
||||
obs = self.normalize_obs(raw_obs)
|
||||
return obs, reward, done, info
|
||||
|
||||
def render(self, **kwargs):
|
||||
return self.env.render()
|
||||
Vendored
+283
@@ -0,0 +1,283 @@
|
||||
"""
|
||||
Multi-step wrapper. Allow executing multiple environmnt steps. Returns stacked observation and optionally stacked previous action.
|
||||
|
||||
Modified from https://github.com/real-stanford/diffusion_policy/blob/main/diffusion_policy/gym_util/multistep_wrapper.py
|
||||
|
||||
"""
|
||||
|
||||
import gym
|
||||
from typing import Optional
|
||||
from gym import spaces
|
||||
import numpy as np
|
||||
from collections import defaultdict, deque
|
||||
|
||||
# import dill
|
||||
|
||||
|
||||
def stack_repeated(x, n):
|
||||
return np.repeat(np.expand_dims(x, axis=0), n, axis=0)
|
||||
|
||||
|
||||
def repeated_box(box_space, n):
|
||||
return spaces.Box(
|
||||
low=stack_repeated(box_space.low, n),
|
||||
high=stack_repeated(box_space.high, n),
|
||||
shape=(n,) + box_space.shape,
|
||||
dtype=box_space.dtype,
|
||||
)
|
||||
|
||||
|
||||
def repeated_space(space, n):
|
||||
if isinstance(space, spaces.Box):
|
||||
return repeated_box(space, n)
|
||||
elif isinstance(space, spaces.Dict):
|
||||
result_space = spaces.Dict()
|
||||
for key, value in space.items():
|
||||
result_space[key] = repeated_space(value, n)
|
||||
return result_space
|
||||
else:
|
||||
raise RuntimeError(f"Unsupported space type {type(space)}")
|
||||
|
||||
|
||||
def take_last_n(x, n):
|
||||
x = list(x)
|
||||
n = min(len(x), n)
|
||||
return np.array(x[-n:])
|
||||
|
||||
|
||||
def dict_take_last_n(x, n):
|
||||
result = dict()
|
||||
for key, value in x.items():
|
||||
result[key] = take_last_n(value, n)
|
||||
return result
|
||||
|
||||
|
||||
def aggregate(data, method="max"):
|
||||
if method == "max":
|
||||
# equivalent to any
|
||||
return np.max(data)
|
||||
elif method == "min":
|
||||
# equivalent to all
|
||||
return np.min(data)
|
||||
elif method == "mean":
|
||||
return np.mean(data)
|
||||
elif method == "sum":
|
||||
return np.sum(data)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
def stack_last_n_obs(all_obs, n_steps):
|
||||
"""Apply padding"""
|
||||
assert len(all_obs) > 0
|
||||
all_obs = list(all_obs)
|
||||
result = np.zeros((n_steps,) + all_obs[-1].shape, dtype=all_obs[-1].dtype)
|
||||
start_idx = -min(n_steps, len(all_obs))
|
||||
result[start_idx:] = np.array(all_obs[start_idx:])
|
||||
if n_steps > len(all_obs):
|
||||
# pad
|
||||
result[:start_idx] = result[start_idx]
|
||||
return result
|
||||
|
||||
|
||||
class MultiStep(gym.Wrapper):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
env,
|
||||
n_obs_steps=1,
|
||||
n_action_steps=1,
|
||||
max_episode_steps=None,
|
||||
reward_agg_method="sum", # never use other types
|
||||
prev_action=True,
|
||||
reset_within_step=False,
|
||||
pass_full_observations=False,
|
||||
verbose=False,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(env)
|
||||
self._single_action_space = env.action_space
|
||||
self._action_space = repeated_space(env.action_space, n_action_steps)
|
||||
self._observation_space = repeated_space(env.observation_space, n_obs_steps)
|
||||
self.max_episode_steps = max_episode_steps
|
||||
self.n_obs_steps = n_obs_steps
|
||||
self.n_action_steps = n_action_steps
|
||||
self.reward_agg_method = reward_agg_method
|
||||
self.prev_action = prev_action
|
||||
self.reset_within_step = reset_within_step
|
||||
self.pass_full_observations = pass_full_observations
|
||||
self.verbose = verbose
|
||||
|
||||
def reset(
|
||||
self,
|
||||
seed: Optional[int] = None,
|
||||
return_info: bool = False,
|
||||
options: dict = {},
|
||||
):
|
||||
"""Resets the environment."""
|
||||
obs = self.env.reset(
|
||||
seed=seed,
|
||||
options=options,
|
||||
return_info=return_info,
|
||||
)
|
||||
self.obs = deque([obs], maxlen=max(self.n_obs_steps + 1, self.n_action_steps))
|
||||
if self.prev_action:
|
||||
self.action = deque(
|
||||
[self._single_action_space.sample()], maxlen=self.n_obs_steps
|
||||
)
|
||||
self.reward = list()
|
||||
self.done = list()
|
||||
self.info = defaultdict(lambda: deque(maxlen=self.n_obs_steps + 1))
|
||||
obs = self._get_obs(self.n_obs_steps)
|
||||
|
||||
self.cnt = 0
|
||||
return obs
|
||||
|
||||
def step(self, action):
|
||||
"""
|
||||
actions: (n_action_steps,) + action_shape
|
||||
"""
|
||||
if action.ndim == 1: # in case action_steps = 1
|
||||
action = action[None]
|
||||
for act_step, act in enumerate(action):
|
||||
self.cnt += 1
|
||||
|
||||
if len(self.done) > 0 and self.done[-1]:
|
||||
# termination
|
||||
break
|
||||
observation, reward, done, info = self.env.step(act)
|
||||
|
||||
self.obs.append(observation)
|
||||
self.action.append(act)
|
||||
self.reward.append(reward)
|
||||
if (
|
||||
self.max_episode_steps is not None
|
||||
) and self.cnt >= self.max_episode_steps:
|
||||
# truncation
|
||||
done = True
|
||||
self.done.append(done)
|
||||
self._add_info(info)
|
||||
|
||||
observation = self._get_obs(self.n_obs_steps)
|
||||
reward = aggregate(self.reward, self.reward_agg_method)
|
||||
done = aggregate(self.done, "max")
|
||||
info = dict_take_last_n(self.info, self.n_obs_steps)
|
||||
if self.pass_full_observations: # right now this assume n_obs_steps = 1
|
||||
info["full_obs"] = self._get_obs(act_step + 1)
|
||||
|
||||
# In mujoco case, done can happen within the loop above
|
||||
if self.reset_within_step and self.done[-1]:
|
||||
observation = (
|
||||
self.reset()
|
||||
) # TODO: arguments? this cannot handle video recording right now since needs to pass in options
|
||||
self.verbose and print("Reset env within wrapper.")
|
||||
|
||||
# reset reward and done for next step
|
||||
self.reward = list()
|
||||
self.done = list()
|
||||
return observation, reward, done, info
|
||||
|
||||
def _get_obs(self, n_steps=1):
|
||||
"""
|
||||
Output (n_steps,) + obs_shape
|
||||
"""
|
||||
assert len(self.obs) > 0
|
||||
if isinstance(self.observation_space, spaces.Box):
|
||||
return stack_last_n_obs(self.obs, n_steps)
|
||||
elif isinstance(self.observation_space, spaces.Dict):
|
||||
result = dict()
|
||||
for key in self.observation_space.keys():
|
||||
result[key] = stack_last_n_obs([obs[key] for obs in self.obs], n_steps)
|
||||
return result
|
||||
else:
|
||||
raise RuntimeError("Unsupported space type")
|
||||
|
||||
def get_prev_action(self, n_steps=None):
|
||||
if n_steps is None:
|
||||
n_steps = self.n_obs_steps - 1 # exclude current step
|
||||
assert len(self.action) > 0
|
||||
return stack_last_n_obs(self.action, n_steps)
|
||||
|
||||
def _add_info(self, info):
|
||||
for key, value in info.items():
|
||||
self.info[key].append(value)
|
||||
|
||||
def render(self, **kwargs):
|
||||
"""Not the best design"""
|
||||
return self.env.render(**kwargs)
|
||||
|
||||
# def get_rewards(self):
|
||||
# return self.reward
|
||||
|
||||
# def get_attr(self, name):
|
||||
# return getattr(self, name)
|
||||
|
||||
# def run_dill_function(self, dill_fn):
|
||||
# fn = dill.loads(dill_fn)
|
||||
# return fn(self)
|
||||
|
||||
# def get_infos(self):
|
||||
# result = dict()
|
||||
# for k, v in self.info.items():
|
||||
# result[k] = list(v)
|
||||
# return result
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import os
|
||||
from omegaconf import OmegaConf
|
||||
import json
|
||||
|
||||
os.environ["MUJOCO_GL"] = "egl"
|
||||
|
||||
cfg = OmegaConf.load("cfg/robomimic/finetune/can/ft_ppo_diffusion_mlp_img.yaml")
|
||||
shape_meta = cfg["shape_meta"]
|
||||
|
||||
import robomimic.utils.env_utils as EnvUtils
|
||||
import robomimic.utils.obs_utils as ObsUtils
|
||||
import matplotlib.pyplot as plt
|
||||
from env.gym_utils.wrapper.robomimic_image import RobomimicImageWrapper
|
||||
|
||||
wrappers = cfg.env.wrappers
|
||||
obs_modality_dict = {
|
||||
"low_dim": (
|
||||
wrappers.robomimic_image.low_dim_keys
|
||||
if "robomimic_image" in wrappers
|
||||
else wrappers.robomimic_lowdim.low_dim_keys
|
||||
),
|
||||
"rgb": (
|
||||
wrappers.robomimic_image.image_keys
|
||||
if "robomimic_image" in wrappers
|
||||
else None
|
||||
),
|
||||
}
|
||||
if obs_modality_dict["rgb"] is None:
|
||||
obs_modality_dict.pop("rgb")
|
||||
ObsUtils.initialize_obs_modality_mapping_from_dict(obs_modality_dict)
|
||||
|
||||
with open(cfg.robomimic_env_cfg_path, "r") as f:
|
||||
env_meta = json.load(f)
|
||||
env = EnvUtils.create_env_from_metadata(
|
||||
env_meta=env_meta,
|
||||
render=False,
|
||||
render_offscreen=False,
|
||||
use_image_obs=True,
|
||||
)
|
||||
env.env.hard_reset = False
|
||||
|
||||
wrapper = MultiStep(
|
||||
env=RobomimicImageWrapper(
|
||||
env=env,
|
||||
shape_meta=shape_meta,
|
||||
image_keys=["robot0_eye_in_hand_image"],
|
||||
),
|
||||
n_obs_steps=1,
|
||||
n_action_steps=1,
|
||||
)
|
||||
wrapper.seed(0)
|
||||
obs = wrapper.reset()
|
||||
print(obs.keys())
|
||||
img = wrapper.render()
|
||||
wrapper.close()
|
||||
plt.imshow(img)
|
||||
plt.savefig("test.png")
|
||||
+227
@@ -0,0 +1,227 @@
|
||||
"""
|
||||
Environment wrapper for Robomimic environments with image observations.
|
||||
|
||||
Modified from https://github.com/real-stanford/diffusion_policy/blob/main/diffusion_policy/env/robomimic/robomimic_image_wrapper.py
|
||||
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import gym
|
||||
from gym import spaces
|
||||
import imageio
|
||||
|
||||
|
||||
class RobomimicImageWrapper(gym.Env):
|
||||
def __init__(
|
||||
self,
|
||||
env,
|
||||
shape_meta: dict,
|
||||
normalization_path=None,
|
||||
low_dim_keys=[
|
||||
"robot0_eef_pos",
|
||||
"robot0_eef_quat",
|
||||
"robot0_gripper_qpos",
|
||||
],
|
||||
image_keys=[
|
||||
"agentview_image",
|
||||
"robot0_eye_in_hand_image",
|
||||
],
|
||||
clamp_obs=False,
|
||||
init_state=None,
|
||||
render_hw=(256, 256),
|
||||
render_camera_name="agentview",
|
||||
):
|
||||
self.env = env
|
||||
self.init_state = init_state
|
||||
self.has_reset_before = False
|
||||
self.render_hw = render_hw
|
||||
self.render_camera_name = render_camera_name
|
||||
self.video_writer = None
|
||||
self.clamp_obs = clamp_obs
|
||||
|
||||
# set up normalization
|
||||
self.normalize = normalization_path is not None
|
||||
if self.normalize:
|
||||
normalization = np.load(normalization_path)
|
||||
self.obs_min = normalization["obs_min"]
|
||||
self.obs_max = normalization["obs_max"]
|
||||
self.action_min = normalization["action_min"]
|
||||
self.action_max = normalization["action_max"]
|
||||
|
||||
# setup spaces
|
||||
low = np.full(env.action_dimension, fill_value=-1)
|
||||
high = np.full(env.action_dimension, fill_value=1)
|
||||
self.action_space = gym.spaces.Box(
|
||||
low=low,
|
||||
high=high,
|
||||
shape=low.shape,
|
||||
dtype=low.dtype,
|
||||
)
|
||||
self.low_dim_keys = low_dim_keys
|
||||
self.image_keys = image_keys
|
||||
self.obs_keys = low_dim_keys + image_keys
|
||||
observation_space = spaces.Dict()
|
||||
for key, value in shape_meta["obs"].items():
|
||||
shape = value["shape"]
|
||||
if key.endswith("rgb"):
|
||||
min_value, max_value = 0, 1
|
||||
elif key.endswith("state"):
|
||||
min_value, max_value = -1, 1
|
||||
else:
|
||||
raise RuntimeError(f"Unsupported type {key}")
|
||||
this_space = spaces.Box(
|
||||
low=min_value,
|
||||
high=max_value,
|
||||
shape=shape,
|
||||
dtype=np.float32,
|
||||
)
|
||||
observation_space[key] = this_space
|
||||
self.observation_space = observation_space
|
||||
|
||||
def normalize_obs(self, obs):
|
||||
obs = 2 * (
|
||||
(obs - self.obs_min) / (self.obs_max - self.obs_min + 1e-6) - 0.5
|
||||
) # -> [-1, 1]
|
||||
if self.clamp_obs:
|
||||
obs = np.clip(obs, -1, 1)
|
||||
return obs
|
||||
|
||||
def unnormalize_action(self, action):
|
||||
action = (action + 1) / 2 # [-1, 1] -> [0, 1]
|
||||
return action * (self.action_max - self.action_min) + self.action_min
|
||||
|
||||
def get_observation(self, raw_obs=None):
|
||||
if raw_obs is None:
|
||||
raw_obs = self.env.get_observation()
|
||||
obs = {"rgb": None, "state": None} # stack rgb if multiple cameras
|
||||
for key in self.obs_keys:
|
||||
if key in self.image_keys:
|
||||
if obs["rgb"] is None:
|
||||
obs["rgb"] = raw_obs[key]
|
||||
else:
|
||||
obs["rgb"] = np.concatenate(
|
||||
[obs["rgb"], raw_obs[key]], axis=0
|
||||
) # C H W
|
||||
else:
|
||||
if obs["state"] is None:
|
||||
obs["state"] = raw_obs[key]
|
||||
else:
|
||||
obs["state"] = np.concatenate([obs["state"], raw_obs[key]], axis=-1)
|
||||
if self.normalize:
|
||||
obs["state"] = self.normalize_obs(obs["state"])
|
||||
obs["rgb"] *= 255 # [0, 1] -> [0, 255], in float64
|
||||
return obs
|
||||
|
||||
def seed(self, seed=None):
|
||||
if seed is not None:
|
||||
np.random.seed(seed=seed)
|
||||
else:
|
||||
np.random.seed()
|
||||
|
||||
def reset(self, options={}, **kwargs):
|
||||
"""Ignore passed-in arguments like seed"""
|
||||
# Close video if exists
|
||||
if self.video_writer is not None:
|
||||
self.video_writer.close()
|
||||
self.video_writer = None
|
||||
|
||||
# Start video if specified
|
||||
if "video_path" in options:
|
||||
self.video_writer = imageio.get_writer(options["video_path"], fps=30)
|
||||
|
||||
# Call reset
|
||||
new_seed = options.get(
|
||||
"seed", None
|
||||
) # used to set all environments to specified seeds
|
||||
if self.init_state is not None:
|
||||
if not self.has_reset_before:
|
||||
# the env must be fully reset at least once to ensure correct rendering
|
||||
self.env.reset()
|
||||
self.has_reset_before = True
|
||||
|
||||
# always reset to the same state to be compatible with gym
|
||||
raw_obs = self.env.reset_to({"states": self.init_state})
|
||||
elif new_seed is not None:
|
||||
self.seed(seed=new_seed)
|
||||
raw_obs = self.env.reset()
|
||||
else:
|
||||
# random reset
|
||||
raw_obs = self.env.reset()
|
||||
return self.get_observation(raw_obs)
|
||||
|
||||
def step(self, action):
|
||||
if self.normalize:
|
||||
action = self.unnormalize_action(action)
|
||||
raw_obs, reward, done, info = self.env.step(action)
|
||||
obs = self.get_observation(raw_obs)
|
||||
|
||||
# render if specified
|
||||
if self.video_writer is not None:
|
||||
video_img = self.render(mode="rgb_array")
|
||||
self.video_writer.append_data(video_img)
|
||||
|
||||
return obs, reward, done, info
|
||||
|
||||
def render(self, mode="rgb_array"):
|
||||
h, w = self.render_hw
|
||||
return self.env.render(
|
||||
mode=mode,
|
||||
height=h,
|
||||
width=w,
|
||||
camera_name=self.render_camera_name,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import os
|
||||
from omegaconf import OmegaConf
|
||||
import json
|
||||
|
||||
os.environ["MUJOCO_GL"] = "egl"
|
||||
|
||||
cfg = OmegaConf.load("cfg/robomimic/finetune/can/ft_ppo_diffusion_mlp_img.yaml")
|
||||
shape_meta = cfg["shape_meta"]
|
||||
|
||||
import robomimic.utils.env_utils as EnvUtils
|
||||
import robomimic.utils.obs_utils as ObsUtils
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
wrappers = cfg.env.wrappers
|
||||
obs_modality_dict = {
|
||||
"low_dim": (
|
||||
wrappers.robomimic_image.low_dim_keys
|
||||
if "robomimic_image" in wrappers
|
||||
else wrappers.robomimic_lowdim.low_dim_keys
|
||||
),
|
||||
"rgb": (
|
||||
wrappers.robomimic_image.image_keys
|
||||
if "robomimic_image" in wrappers
|
||||
else None
|
||||
),
|
||||
}
|
||||
if obs_modality_dict["rgb"] is None:
|
||||
obs_modality_dict.pop("rgb")
|
||||
ObsUtils.initialize_obs_modality_mapping_from_dict(obs_modality_dict)
|
||||
|
||||
with open(cfg.robomimic_env_cfg_path, "r") as f:
|
||||
env_meta = json.load(f)
|
||||
env = EnvUtils.create_env_from_metadata(
|
||||
env_meta=env_meta,
|
||||
render=False,
|
||||
render_offscreen=False,
|
||||
use_image_obs=True,
|
||||
)
|
||||
env.env.hard_reset = False
|
||||
|
||||
wrapper = RobomimicImageWrapper(
|
||||
env=env,
|
||||
shape_meta=shape_meta,
|
||||
image_keys=["robot0_eye_in_hand_image"],
|
||||
)
|
||||
wrapper.seed(0)
|
||||
obs = wrapper.reset()
|
||||
print(obs.keys())
|
||||
img = wrapper.render()
|
||||
wrapper.close()
|
||||
plt.imshow(img)
|
||||
plt.savefig("test.png")
|
||||
+142
@@ -0,0 +1,142 @@
|
||||
"""
|
||||
Environment wrapper for Robomimic environments with state observations.
|
||||
|
||||
Modified from https://github.com/real-stanford/diffusion_policy/blob/main/diffusion_policy/env/robomimic/robomimic_lowdim_wrapper.py
|
||||
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import gym
|
||||
from gym.spaces import Box
|
||||
import imageio
|
||||
|
||||
|
||||
class RobomimicLowdimWrapper(gym.Env):
|
||||
def __init__(
|
||||
self,
|
||||
env,
|
||||
normalization_path=None,
|
||||
low_dim_keys=[
|
||||
"robot0_eef_pos",
|
||||
"robot0_eef_quat",
|
||||
"robot0_gripper_qpos",
|
||||
"object",
|
||||
],
|
||||
clamp_obs=False,
|
||||
init_state=None,
|
||||
render_hw=(256, 256),
|
||||
render_camera_name="agentview",
|
||||
):
|
||||
self.env = env
|
||||
self.obs_keys = low_dim_keys
|
||||
self.init_state = init_state
|
||||
self.render_hw = render_hw
|
||||
self.render_camera_name = render_camera_name
|
||||
self.video_writer = None
|
||||
self.clamp_obs = clamp_obs
|
||||
|
||||
# set up normalization
|
||||
self.normalize = normalization_path is not None
|
||||
if self.normalize:
|
||||
normalization = np.load(normalization_path)
|
||||
self.obs_min = normalization["obs_min"]
|
||||
self.obs_max = normalization["obs_max"]
|
||||
self.action_min = normalization["action_min"]
|
||||
self.action_max = normalization["action_max"]
|
||||
|
||||
# setup spaces - use [-1, 1]
|
||||
low = np.full(env.action_dimension, fill_value=-1)
|
||||
high = np.full(env.action_dimension, fill_value=1)
|
||||
self.action_space = Box(
|
||||
low=low,
|
||||
high=high,
|
||||
shape=low.shape,
|
||||
dtype=low.dtype,
|
||||
)
|
||||
obs_example = self.get_observation()
|
||||
low = np.full_like(obs_example, fill_value=-1)
|
||||
high = np.full_like(obs_example, fill_value=1)
|
||||
self.observation_space = Box(
|
||||
low=low,
|
||||
high=high,
|
||||
shape=low.shape,
|
||||
dtype=low.dtype,
|
||||
)
|
||||
|
||||
def normalize_obs(self, obs):
|
||||
obs = 2 * (
|
||||
(obs - self.obs_min) / (self.obs_max - self.obs_min + 1e-6) - 0.5
|
||||
) # -> [-1, 1]
|
||||
if self.clamp_obs:
|
||||
obs = np.clip(obs, -1, 1)
|
||||
return obs
|
||||
|
||||
def unnormalize_action(self, action):
|
||||
action = (action + 1) / 2 # [-1, 1] -> [0, 1]
|
||||
return action * (self.action_max - self.action_min) + self.action_min
|
||||
|
||||
def get_observation(self):
|
||||
raw_obs = self.env.get_observation()
|
||||
raw_obs = np.concatenate([raw_obs[key] for key in self.obs_keys], axis=0)
|
||||
if self.normalize:
|
||||
return self.normalize_obs(raw_obs)
|
||||
return raw_obs
|
||||
|
||||
def seed(self, seed=None):
|
||||
if seed is not None:
|
||||
np.random.seed(seed=seed)
|
||||
else:
|
||||
np.random.seed()
|
||||
|
||||
def reset(self, options={}, **kwargs):
|
||||
"""Ignore passed-in arguments like seed"""
|
||||
|
||||
# Close video if exists
|
||||
if self.video_writer is not None:
|
||||
self.video_writer.close()
|
||||
self.video_writer = None
|
||||
|
||||
# Start video if specified
|
||||
if "video_path" in options:
|
||||
self.video_writer = imageio.get_writer(options["video_path"], fps=30)
|
||||
|
||||
# Call reset
|
||||
new_seed = options.get(
|
||||
"seed", None
|
||||
) # used to set all environments to specified seeds
|
||||
if self.init_state is not None:
|
||||
# always reset to the same state to be compatible with gym
|
||||
self.env.reset_to({"states": self.init_state})
|
||||
elif new_seed is not None:
|
||||
self.seed(seed=new_seed)
|
||||
self.env.reset()
|
||||
else:
|
||||
# random reset
|
||||
self.env.reset()
|
||||
return self.get_observation()
|
||||
|
||||
def step(self, action):
|
||||
if self.normalize:
|
||||
action = self.unnormalize_action(action)
|
||||
raw_obs, reward, done, info = self.env.step(action)
|
||||
raw_obs = np.concatenate([raw_obs[key] for key in self.obs_keys], axis=0)
|
||||
if self.normalize:
|
||||
obs = self.normalize_obs(raw_obs)
|
||||
else:
|
||||
obs = raw_obs
|
||||
|
||||
# render if specified
|
||||
if self.video_writer is not None:
|
||||
video_img = self.render(mode="rgb_array")
|
||||
self.video_writer.append_data(video_img)
|
||||
|
||||
return obs, reward, done, info
|
||||
|
||||
def render(self, mode="rgb_array"):
|
||||
h, w = self.render_hw
|
||||
return self.env.render(
|
||||
mode=mode,
|
||||
height=h,
|
||||
width=w,
|
||||
camera_name=self.render_camera_name,
|
||||
)
|
||||
Reference in New Issue
Block a user