Fixes build errors due to name conflicts
This commit is contained in:
@@ -0,0 +1,375 @@
|
||||
from functools import partial
|
||||
from typing import Any, Tuple, Union
|
||||
|
||||
import chex
|
||||
import gymnax
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
from brax import envs
|
||||
from brax.envs.wrappers.training import AutoResetWrapper, EpisodeWrapper
|
||||
from flax import struct
|
||||
from gymnax.environments import environment, spaces
|
||||
from gymnax.environments.environment import Environment
|
||||
from gymnax.environments.spaces import Box
|
||||
from ml_collections import ConfigDict
|
||||
from mujoco_playground import MjxEnv, registry
|
||||
from mujoco_playground._src.wrapper import wrap_for_brax_training, Wrapper
|
||||
|
||||
|
||||
class MjxGymnaxWrapper(Environment):
|
||||
def __init__(
|
||||
self,
|
||||
env_or_name: str | MjxEnv,
|
||||
episode_length: int = 1000,
|
||||
action_repeat: int = 1,
|
||||
reward_scale: float = 1.0,
|
||||
push_distractions: bool = False,
|
||||
config: dict = None,
|
||||
asymmetric_observation: bool = False,
|
||||
):
|
||||
if isinstance(env_or_name, str):
|
||||
if config is None:
|
||||
config = registry.get_default_config(env_or_name)
|
||||
is_humanoid_task = env_or_name in [
|
||||
"G1JoystickRoughTerrain",
|
||||
"G1JoystickFlatTerrain",
|
||||
"T1JoystickRoughTerrain",
|
||||
"T1JoystickFlatTerrain",
|
||||
]
|
||||
if is_humanoid_task:
|
||||
config.push_config.enable = push_distractions
|
||||
else:
|
||||
config = ConfigDict(config)
|
||||
env = registry.load(env_or_name, config=config)
|
||||
if episode_length is not None:
|
||||
env = wrap_for_brax_training(
|
||||
env, episode_length=episode_length, action_repeat=action_repeat
|
||||
)
|
||||
self.env = env
|
||||
else:
|
||||
self.env = env_or_name
|
||||
self.reward_scale = reward_scale
|
||||
if isinstance(self.env.observation_size, int):
|
||||
self.dict_obs = False
|
||||
else:
|
||||
self.dict_obs = True
|
||||
if asymmetric_observation:
|
||||
self.dict_obs_key = "privileged_state"
|
||||
else:
|
||||
self.dict_obs_key = "state"
|
||||
print(self.dict_obs_key)
|
||||
super().__init__()
|
||||
|
||||
def action_space(self, params):
|
||||
return gymnax.environments.spaces.Box(
|
||||
low=-1.0,
|
||||
high=1.0,
|
||||
shape=(self.env.action_size,),
|
||||
)
|
||||
|
||||
def observation_space(self, params):
|
||||
if self.dict_obs:
|
||||
return Box(
|
||||
low=-float("inf"),
|
||||
high=float("inf"),
|
||||
shape=self.env.observation_size["state"],
|
||||
), Box(
|
||||
low=-float("inf"),
|
||||
high=float("inf"),
|
||||
shape=self.env.observation_size[self.dict_obs_key],
|
||||
)
|
||||
else:
|
||||
return Box(
|
||||
low=-float("inf"),
|
||||
high=float("inf"),
|
||||
shape=(self.env.observation_size,),
|
||||
), Box(
|
||||
low=-float("inf"),
|
||||
high=float("inf"),
|
||||
shape=(self.env.observation_size,),
|
||||
)
|
||||
|
||||
@property
|
||||
def default_params(self) -> gymnax.EnvParams:
|
||||
return gymnax.EnvParams()
|
||||
|
||||
def reset(self, key):
|
||||
state = self.env.reset(key)
|
||||
# state.info["truncation"] = 0.0
|
||||
obs = state.obs if not self.dict_obs else state.obs["state"]
|
||||
critic_obs = state.obs if not self.dict_obs else state.obs[self.dict_obs_key]
|
||||
return obs, critic_obs, state
|
||||
|
||||
def step(self, key, state, action):
|
||||
# action = jnp.nan_to_num(action, 0.0)
|
||||
state = self.env.step(state, action)
|
||||
obs = state.obs if not self.dict_obs else state.obs["state"]
|
||||
critic_obs = state.obs if not self.dict_obs else state.obs[self.dict_obs_key]
|
||||
return (
|
||||
obs,
|
||||
critic_obs,
|
||||
state,
|
||||
state.reward * self.reward_scale,
|
||||
state.done > 0.5,
|
||||
{},
|
||||
)
|
||||
|
||||
|
||||
@struct.dataclass
|
||||
class LogEnvState:
|
||||
env_state: environment.EnvState
|
||||
episode_returns: jnp.ndarray
|
||||
episode_lengths: jnp.ndarray
|
||||
returned_episode_returns: jnp.ndarray
|
||||
returned_episode_lengths: jnp.ndarray
|
||||
timestep: jnp.ndarray
|
||||
truncated: jnp.ndarray
|
||||
info: Any = None
|
||||
|
||||
def unwrapped(self):
|
||||
return self.env_state
|
||||
|
||||
def set_env_state(self, env_state):
|
||||
return self.replace(env_state=env_state)
|
||||
|
||||
|
||||
class LogWrapper(Wrapper):
|
||||
"""Log the episode returns and lengths."""
|
||||
|
||||
def __init__(self, env: environment.Environment, num_envs: int):
|
||||
super().__init__(env)
|
||||
self.num_envs = num_envs
|
||||
|
||||
@partial(jax.jit, static_argnums=(0,))
|
||||
def reset(self, key) -> Tuple[chex.Array, environment.EnvState]:
|
||||
obs, critic_obs, env_state = self.env.reset(key)
|
||||
state = LogEnvState(
|
||||
env_state=env_state,
|
||||
episode_returns=jnp.zeros((self.num_envs,)),
|
||||
episode_lengths=jnp.zeros((self.num_envs,), dtype=jnp.int32),
|
||||
returned_episode_returns=jnp.zeros((self.num_envs,)),
|
||||
returned_episode_lengths=jnp.zeros((self.num_envs,), dtype=jnp.int32),
|
||||
timestep=jnp.zeros((self.num_envs,), dtype=jnp.int32),
|
||||
truncated=jnp.ones((self.num_envs,), dtype=jnp.float32),
|
||||
info={
|
||||
"returned_episode": jnp.zeros((self.num_envs,), dtype=jnp.bool_),
|
||||
"returned_episode_returns": jnp.zeros((self.num_envs,)),
|
||||
"timestep": jnp.zeros((self.num_envs,), dtype=jnp.int32),
|
||||
"returned_episode_lengths": jnp.zeros(
|
||||
(self.num_envs,), dtype=jnp.int32
|
||||
),
|
||||
},
|
||||
)
|
||||
return obs, critic_obs, state
|
||||
|
||||
@partial(jax.jit, static_argnums=(0,))
|
||||
def step(
|
||||
self,
|
||||
key: chex.PRNGKey,
|
||||
state: environment.EnvState,
|
||||
action: Union[int, float],
|
||||
) -> Tuple[chex.Array, environment.EnvState, float, bool, dict]:
|
||||
obs, critic_obs, env_state, reward, done, info = self.env.step(
|
||||
key, state.env_state, action
|
||||
)
|
||||
new_episode_return = state.episode_returns + reward
|
||||
new_episode_length = state.episode_lengths + 1
|
||||
info["returned_episode_returns"] = (
|
||||
state.returned_episode_returns * (1 - done) + new_episode_return * done
|
||||
)
|
||||
info["returned_episode_lengths"] = (
|
||||
state.returned_episode_lengths * (1 - done) + new_episode_length * done
|
||||
)
|
||||
info["timestep"] = state.timestep
|
||||
info["returned_episode"] = done
|
||||
state = LogEnvState(
|
||||
env_state=env_state,
|
||||
episode_returns=new_episode_return * (1 - done),
|
||||
episode_lengths=new_episode_length * (1 - done),
|
||||
returned_episode_returns=state.returned_episode_returns * (1 - done)
|
||||
+ new_episode_return * done,
|
||||
returned_episode_lengths=state.returned_episode_lengths * (1 - done)
|
||||
+ new_episode_length * done,
|
||||
timestep=state.timestep + 1,
|
||||
truncated=env_state.info["truncation"],
|
||||
info=info,
|
||||
)
|
||||
return obs, critic_obs, state, reward, done, info
|
||||
|
||||
|
||||
class BraxGymnaxWrapper:
|
||||
def __init__(
|
||||
self,
|
||||
env_name,
|
||||
backend="generalized",
|
||||
episode_length=1000,
|
||||
reward_scaling=1.0,
|
||||
terminate=True,
|
||||
):
|
||||
env = envs.get_environment(
|
||||
env_name=env_name, backend=backend, terminate_when_unhealthy=terminate
|
||||
)
|
||||
env = EpisodeWrapper(env, episode_length=episode_length, action_repeat=1)
|
||||
env = AutoResetWrapper(env)
|
||||
self.env = env
|
||||
self.action_size = self.env.action_size
|
||||
self.observation_size = (self.env.observation_size,)
|
||||
self.default_params = ()
|
||||
self.reward_scaling = reward_scaling
|
||||
|
||||
def reset(self, key):
|
||||
state = self.env.reset(key)
|
||||
return state.obs, state
|
||||
|
||||
def step(self, key, state, action):
|
||||
next_state = self.env.step(state, action)
|
||||
return (
|
||||
next_state.obs,
|
||||
next_state.obs,
|
||||
next_state,
|
||||
next_state.reward * self.reward_scaling,
|
||||
next_state.done > 0.5,
|
||||
{},
|
||||
)
|
||||
|
||||
def observation_space(self):
|
||||
return spaces.Box(
|
||||
low=-jnp.inf,
|
||||
high=jnp.inf,
|
||||
shape=(self.env.observation_size,),
|
||||
), spaces.Box(
|
||||
low=-jnp.inf,
|
||||
high=jnp.inf,
|
||||
shape=(self.env.observation_size,),
|
||||
)
|
||||
|
||||
def action_space(self):
|
||||
return spaces.Box(
|
||||
low=-1.0,
|
||||
high=1.0,
|
||||
shape=(self.env.action_size,),
|
||||
)
|
||||
|
||||
|
||||
class ClipAction(Wrapper):
|
||||
def __init__(self, env, low=-0.999, high=0.999):
|
||||
super().__init__(env)
|
||||
self.low = low
|
||||
self.high = high
|
||||
|
||||
def step(self, key, state, action):
|
||||
"""TODO: In theory the below line should be the way to do this."""
|
||||
# action = jnp.clip(action, self.env.action_space.low, self.env.action_space.high)
|
||||
action = jnp.clip(action, self.low, self.high)
|
||||
return self.env.step(key, state, action)
|
||||
|
||||
|
||||
@struct.dataclass
|
||||
class NormalizeVecObsEnvState:
|
||||
mean: jnp.ndarray
|
||||
var: jnp.ndarray
|
||||
critic_mean: jnp.ndarray
|
||||
critic_var: jnp.ndarray
|
||||
count: float
|
||||
env_state: environment.EnvState
|
||||
truncated: float
|
||||
info: Any = None
|
||||
|
||||
def unwrapped(self):
|
||||
return self.env_state.unwrapped()
|
||||
|
||||
def set_env_state(self, env_state):
|
||||
return self.replace(env_state=self.env_state.set_env_state(env_state))
|
||||
|
||||
|
||||
class NormalizeVec(Wrapper):
|
||||
def __init__(self, env):
|
||||
super().__init__(env)
|
||||
|
||||
def _init_state(self, key):
|
||||
obs, critic_obs, env_state = self.env.reset(key)
|
||||
return NormalizeVecObsEnvState(
|
||||
mean=jnp.mean(obs, axis=0),
|
||||
var=jnp.var(obs, axis=0),
|
||||
critic_mean=jnp.mean(critic_obs, axis=0),
|
||||
critic_var=jnp.var(critic_obs, axis=0),
|
||||
count=obs.shape[0],
|
||||
env_state=env_state,
|
||||
)
|
||||
|
||||
def _compute_stats(self, mean, var, count, obs):
|
||||
batch_mean = jnp.mean(obs, axis=0)
|
||||
batch_var = jnp.var(obs, axis=0)
|
||||
batch_count = obs.shape[0]
|
||||
|
||||
delta = batch_mean - mean
|
||||
tot_count = count + batch_count
|
||||
|
||||
new_mean = mean + delta * batch_count / tot_count
|
||||
m_a = var * count
|
||||
m_b = batch_var * batch_count
|
||||
M2 = m_a + m_b + jnp.square(delta) * count * batch_count / tot_count
|
||||
new_var = M2 / tot_count
|
||||
|
||||
return new_mean, new_var
|
||||
|
||||
def reset(self, key, params=None):
|
||||
obs, critic_obs, env_state = self.env.reset(key)
|
||||
if params is not None:
|
||||
mean = params.mean
|
||||
var = params.var
|
||||
critic_mean = params.critic_mean
|
||||
critic_var = params.critic_var
|
||||
count = params.count
|
||||
else:
|
||||
mean = jnp.mean(obs, axis=0)
|
||||
var = jnp.var(obs, axis=0)
|
||||
critic_mean = jnp.mean(critic_obs, axis=0)
|
||||
critic_var = jnp.var(critic_obs, axis=0)
|
||||
count = obs.shape[0]
|
||||
state = NormalizeVecObsEnvState(
|
||||
mean=mean,
|
||||
var=var,
|
||||
critic_mean=critic_mean,
|
||||
critic_var=critic_var,
|
||||
count=count,
|
||||
env_state=env_state,
|
||||
truncated=env_state.truncated,
|
||||
info=env_state.info,
|
||||
)
|
||||
return (
|
||||
(obs - state.mean) / jnp.sqrt(state.var + 1e-2),
|
||||
(critic_obs - state.critic_mean) / jnp.sqrt(state.critic_var + 1e-2),
|
||||
state,
|
||||
)
|
||||
|
||||
def step(self, key, state, action):
|
||||
obs, critic_obs, env_state, reward, done, info = self.env.step(
|
||||
key, state.env_state, action
|
||||
)
|
||||
|
||||
new_mean, new_var = self._compute_stats(state.mean, state.var, state.count, obs)
|
||||
new_critic_mean, new_critic_var = self._compute_stats(
|
||||
state.critic_mean, state.critic_var, state.count, critic_obs
|
||||
)
|
||||
|
||||
new_count = state.count + obs.shape[0]
|
||||
|
||||
state = NormalizeVecObsEnvState(
|
||||
mean=new_mean,
|
||||
var=new_var,
|
||||
critic_mean=new_critic_mean,
|
||||
critic_var=new_critic_var,
|
||||
count=new_count,
|
||||
env_state=env_state,
|
||||
truncated=env_state.truncated,
|
||||
info=env_state.info,
|
||||
)
|
||||
return (
|
||||
(obs - state.mean) / jnp.sqrt(state.var + 1e-2),
|
||||
(critic_obs - state.critic_mean) / jnp.sqrt(state.critic_var + 1e-2),
|
||||
state,
|
||||
reward,
|
||||
done,
|
||||
info,
|
||||
)
|
||||
@@ -0,0 +1,124 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import gymnasium as gym
|
||||
import numpy as np
|
||||
import torch
|
||||
from gymnasium.wrappers import TimeLimit
|
||||
from loguru import logger as log
|
||||
from stable_baselines3.common.vec_env import SubprocVecEnv
|
||||
|
||||
# Disable all logging below CRITICAL level
|
||||
log.remove()
|
||||
log.add(lambda msg: False, level="CRITICAL")
|
||||
|
||||
|
||||
def make_env(env_name, rank, render_mode=None, seed=0):
|
||||
"""
|
||||
Utility function for multiprocessed env.
|
||||
|
||||
:param rank: (int) index of the subprocess
|
||||
:param seed: (int) the inital seed for RNG
|
||||
"""
|
||||
|
||||
if env_name in [
|
||||
"h1hand-push-v0",
|
||||
"h1-push-v0",
|
||||
"h1hand-cube-v0",
|
||||
"h1cube-v0",
|
||||
"h1hand-basketball-v0",
|
||||
"h1-basketball-v0",
|
||||
"h1hand-kitchen-v0",
|
||||
"h1-kitchen-v0",
|
||||
]:
|
||||
max_episode_steps = 500
|
||||
else:
|
||||
max_episode_steps = 1000
|
||||
|
||||
def _init():
|
||||
env = gym.make(env_name, render_mode=render_mode)
|
||||
env = TimeLimit(env, max_episode_steps=max_episode_steps)
|
||||
env.unwrapped.seed(seed + rank)
|
||||
|
||||
return env
|
||||
|
||||
return _init
|
||||
|
||||
|
||||
class HumanoidBenchEnv:
|
||||
"""Wraps HumanoidBench environment to support parallel environments."""
|
||||
|
||||
def __init__(self, env_name, num_envs=1, render_mode=None, device=None):
|
||||
# NOTE: HumanoidBench action space is already normalized to [-1, 1]
|
||||
device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
self.sim_device = device
|
||||
self.num_envs = num_envs
|
||||
|
||||
# Create the base environment
|
||||
self.envs = SubprocVecEnv(
|
||||
[make_env(env_name, i, render_mode=render_mode) for i in range(num_envs)]
|
||||
)
|
||||
|
||||
if env_name in [
|
||||
"h1hand-push-v0",
|
||||
"h1-push-v0",
|
||||
"h1hand-cube-v0",
|
||||
"h1cube-v0",
|
||||
"h1hand-basketball-v0",
|
||||
"h1-basketball-v0",
|
||||
"h1hand-kitchen-v0",
|
||||
"h1-kitchen-v0",
|
||||
]:
|
||||
self.max_episode_steps = 500
|
||||
else:
|
||||
self.max_episode_steps = 1000
|
||||
|
||||
# For compatibility with MuJoCo Playground
|
||||
self.asymmetric_obs = False # For comptatibility with MuJoCo Playground
|
||||
self.num_obs = self.envs.observation_space.shape[-1]
|
||||
self.num_actions = self.envs.action_space.shape[-1]
|
||||
|
||||
def reset(self):
|
||||
"""Reset the environment."""
|
||||
observations = self.envs.reset()
|
||||
observations = torch.from_numpy(observations).to(
|
||||
device=self.sim_device, dtype=torch.float
|
||||
)
|
||||
return observations
|
||||
|
||||
def render(self):
|
||||
assert self.num_envs == 1, (
|
||||
"Currently only supports single environment rendering"
|
||||
)
|
||||
return self.envs.render()
|
||||
|
||||
def step(self, actions):
|
||||
assert isinstance(actions, torch.Tensor)
|
||||
actions = actions.cpu().numpy()
|
||||
|
||||
observations, rewards, dones, raw_infos = self.envs.step(actions)
|
||||
|
||||
# This will be used for getting 'true' next observations
|
||||
infos = dict()
|
||||
infos["observations"] = {"raw": {"obs": observations.copy()}}
|
||||
truncateds = np.zeros_like(dones)
|
||||
for i in range(self.num_envs):
|
||||
if raw_infos[i].get("TimeLimit.truncated", False):
|
||||
truncateds[i] = True
|
||||
infos["observations"]["raw"]["obs"][i] = raw_infos[i][
|
||||
"terminal_observation"
|
||||
]
|
||||
|
||||
observations = torch.from_numpy(observations).to(
|
||||
device=self.sim_device, dtype=torch.float
|
||||
)
|
||||
rewards = torch.from_numpy(rewards).to(
|
||||
device=self.sim_device, dtype=torch.float
|
||||
)
|
||||
dones = torch.from_numpy(dones).to(device=self.sim_device)
|
||||
truncateds = torch.from_numpy(truncateds).to(device=self.sim_device)
|
||||
infos["observations"]["raw"]["obs"] = torch.from_numpy(
|
||||
infos["observations"]["raw"]["obs"]
|
||||
).to(device=self.sim_device, dtype=torch.float)
|
||||
infos["time_outs"] = truncateds
|
||||
|
||||
return observations, rewards, dones, infos
|
||||
@@ -0,0 +1,81 @@
|
||||
from typing import Optional
|
||||
|
||||
import gymnasium as gym
|
||||
import torch
|
||||
from isaaclab.app import AppLauncher
|
||||
from isaaclab_tasks.utils.parse_cfg import parse_env_cfg
|
||||
|
||||
app_launcher = AppLauncher(headless=True)
|
||||
simulation_app = app_launcher.app
|
||||
|
||||
|
||||
|
||||
class IsaacLabEnv:
|
||||
"""Wrapper for IsaacLab environments to be compatible with MuJoCo Playground"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
task_name: str,
|
||||
device: str,
|
||||
num_envs: int,
|
||||
seed: int,
|
||||
action_bounds: Optional[float] = None,
|
||||
):
|
||||
env_cfg = parse_env_cfg(
|
||||
task_name,
|
||||
device=device,
|
||||
num_envs=num_envs,
|
||||
)
|
||||
env_cfg.seed = seed
|
||||
self.seed = seed
|
||||
self.envs = gym.make(task_name, cfg=env_cfg, render_mode=None)
|
||||
|
||||
self.num_envs = self.envs.unwrapped.num_envs
|
||||
self.max_episode_steps = self.envs.unwrapped.max_episode_length
|
||||
self.action_bounds = action_bounds
|
||||
self.num_obs = self.envs.unwrapped.single_observation_space["policy"].shape[0]
|
||||
self.asymmetric_obs = "critic" in self.envs.unwrapped.single_observation_space
|
||||
if self.asymmetric_obs:
|
||||
self.num_privileged_obs = self.envs.unwrapped.single_observation_space[
|
||||
"critic"
|
||||
].shape[0]
|
||||
else:
|
||||
self.num_privileged_obs = 0
|
||||
self.num_actions = self.envs.unwrapped.single_action_space.shape[0]
|
||||
|
||||
def reset(self, random_start_init: bool = True) -> torch.Tensor:
|
||||
obs_dict, _ = self.envs.reset()
|
||||
# NOTE: decorrelate episode horizons like RSL‑RL
|
||||
if random_start_init:
|
||||
self.envs.unwrapped.episode_length_buf = torch.randint_like(
|
||||
self.envs.unwrapped.episode_length_buf, high=int(self.max_episode_steps)
|
||||
)
|
||||
return obs_dict["policy"]
|
||||
|
||||
def reset_with_critic_obs(self) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
obs_dict, _ = self.envs.reset()
|
||||
return obs_dict["policy"], obs_dict["critic"]
|
||||
|
||||
def step(
|
||||
self, actions: torch.Tensor
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict]:
|
||||
if self.action_bounds is not None:
|
||||
actions = torch.clamp(actions, -1.0, 1.0) * self.action_bounds
|
||||
obs_dict, rew, terminations, truncations, infos = self.envs.step(actions)
|
||||
dones = (terminations | truncations).to(dtype=torch.long)
|
||||
obs = obs_dict["policy"]
|
||||
critic_obs = obs_dict["critic"] if self.asymmetric_obs else None
|
||||
info_ret = {"time_outs": truncations, "observations": {"critic": critic_obs}}
|
||||
# NOTE: There's really no way to get the raw observations from IsaacLab
|
||||
# We just use the 'reset_obs' as next_obs, unfortunately.
|
||||
# See https://github.com/isaac-sim/IsaacLab/issues/1362
|
||||
info_ret["observations"]["raw"] = {
|
||||
"obs": obs,
|
||||
"critic_obs": critic_obs,
|
||||
}
|
||||
return obs, rew, dones, info_ret
|
||||
|
||||
def render(self):
|
||||
raise NotImplementedError(
|
||||
"We don't support rendering for IsaacLab environments"
|
||||
)
|
||||
@@ -0,0 +1,89 @@
|
||||
from gymnasium import Wrapper
|
||||
import torch
|
||||
|
||||
|
||||
class ManiSkillWrapper(Wrapper):
|
||||
"""
|
||||
A wrapper for ManiSkill environments to ensure compatibility with the expected API.
|
||||
This wrapper is used to handle the ManiSkill environments in a way that is consistent
|
||||
with the other environments in the codebase.
|
||||
"""
|
||||
|
||||
def __init__(self, env, max_episode_steps: int, partial_reset, device: str):
|
||||
super().__init__(env)
|
||||
self.action_space = env.action_space
|
||||
self.observation_space = env.observation_space
|
||||
self.metadata = env.metadata
|
||||
self.asymmetric_obs = False
|
||||
self.max_episode_steps = max_episode_steps
|
||||
|
||||
self.partial_reset = partial_reset
|
||||
|
||||
self.returns = torch.zeros(env.num_envs, dtype=torch.float32, device=device)
|
||||
self.episode_len = torch.zeros(env.num_envs, dtype=torch.float32, device=device)
|
||||
self.success = torch.zeros(env.num_envs, dtype=torch.float32, device=device)
|
||||
|
||||
@property
|
||||
def unwrapped(self):
|
||||
"""
|
||||
Returns the underlying environment.
|
||||
"""
|
||||
return self.env
|
||||
|
||||
@property
|
||||
def num_actions(self):
|
||||
"""
|
||||
Returns the number of actions in the action space.
|
||||
"""
|
||||
return self.action_space.shape[1]
|
||||
|
||||
@property
|
||||
def num_obs(self):
|
||||
"""
|
||||
Returns the number of observations in the observation space.
|
||||
"""
|
||||
return self.observation_space.shape[1]
|
||||
|
||||
def reset(self, seed=None, options=dict()):
|
||||
"""
|
||||
Resets the environment and returns the initial observation.
|
||||
"""
|
||||
return self.env.reset(seed=seed, options=options)
|
||||
|
||||
def step(self, action):
|
||||
"""
|
||||
Takes a step in the environment with the given action.
|
||||
Returns the next observation, reward, done, and info.
|
||||
"""
|
||||
obs, reward, terminated, truncated, info = self.env.step(action)
|
||||
if "final_info" in info:
|
||||
self.returns = (
|
||||
info["final_info"]["episode"]["return"] * info["_final_info"].float()
|
||||
+ (1.0 - info["_final_info"].float()) * self.returns
|
||||
)
|
||||
self.episode_len = (
|
||||
info["final_info"]["episode"]["episode_len"]
|
||||
* info["_final_info"].float()
|
||||
+ (1.0 - info["_final_info"].float()) * self.episode_len
|
||||
)
|
||||
self.success = (
|
||||
info["final_info"]["episode"]["success_once"]
|
||||
* info["_final_info"].float()
|
||||
+ (1.0 - info["_final_info"].float()) * self.success
|
||||
)
|
||||
info["log_info"] = {
|
||||
"return": self.returns,
|
||||
"episode_len": self.episode_len,
|
||||
"success": self.success,
|
||||
}
|
||||
if self.partial_reset:
|
||||
# maniskill continues bootstrap on terminated, which playground does on truncated.
|
||||
# This unifies the interfaces in a very hacky way
|
||||
done = torch.zeros_like(
|
||||
terminated, dtype=torch.bool, device=terminated.device
|
||||
)
|
||||
truncated = torch.logical_or(terminated, truncated)
|
||||
else:
|
||||
done = torch.logical_or(terminated, truncated)
|
||||
truncated = torch.zeros_like(done, dtype=torch.bool, device=done.device)
|
||||
return obs, reward, done, truncated, info
|
||||
@@ -0,0 +1,148 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import isaacgymenvs
|
||||
import torch
|
||||
from omegaconf import OmegaConf
|
||||
|
||||
|
||||
class MTBenchEnv:
|
||||
def __init__(
|
||||
self,
|
||||
task_name: str,
|
||||
device_id: int,
|
||||
num_envs: int,
|
||||
seed: int,
|
||||
):
|
||||
# NOTE: Currently, we only support Meta-World-v2 MT-10/MT-50 in MTBench
|
||||
task_config = MTBENCH_MW2_CONFIG.copy()
|
||||
if task_name == "meta-world-v2-mt10":
|
||||
# MT-10 Setup
|
||||
assert num_envs == 4096, "MT-10 only supports 4096 environments (for now)"
|
||||
self.num_tasks = 10
|
||||
task_config["env"]["tasks"] = [4, 16, 17, 18, 28, 31, 38, 40, 48, 49]
|
||||
task_config["env"]["taskEnvCount"] = [410] * 6 + [409] * 4
|
||||
elif task_name == "meta-world-v2-mt50":
|
||||
# MT-50 Setup
|
||||
self.num_tasks = 50
|
||||
assert num_envs == 8192, "MT-50 only supports 8192 environments (for now)"
|
||||
task_config["env"]["tasks"] = list(range(50))
|
||||
task_config["env"]["taskEnvCount"] = [164] * 42 + [163] * 8 # 6888 + 1304
|
||||
else:
|
||||
raise ValueError(f"Unsupported task name: {task_name}")
|
||||
task_config["env"]["numEnvs"] = num_envs
|
||||
task_config["env"]["numObservations"] = 39 + self.num_tasks
|
||||
task_config["env"]["seed"] = seed
|
||||
|
||||
# Convert dictionary to OmegaConf object
|
||||
env_cfg = {"task": task_config}
|
||||
env_cfg = OmegaConf.create(env_cfg)
|
||||
|
||||
self.env = isaacgymenvs.make(
|
||||
task=env_cfg.task.name,
|
||||
num_envs=num_envs,
|
||||
sim_device=f"cuda:{device_id}",
|
||||
rl_device=f"cuda:{device_id}",
|
||||
seed=seed,
|
||||
headless=True,
|
||||
cfg=env_cfg,
|
||||
)
|
||||
|
||||
self.num_envs = num_envs
|
||||
self.asymmetric_obs = False
|
||||
self.num_obs = self.env.observation_space.shape[0]
|
||||
assert self.num_obs == 39 + self.num_tasks, (
|
||||
"MTBench observation space is 39 + num_tasks (one-hot vector)"
|
||||
)
|
||||
self.num_privileged_obs = 0
|
||||
self.num_actions = self.env.action_space.shape[0]
|
||||
self.max_episode_steps = self.env.max_episode_length
|
||||
|
||||
def reset(self) -> torch.Tensor:
|
||||
"""Reset the environment."""
|
||||
# TODO: Check if we need no_grad and detach here
|
||||
with torch.no_grad(): # do we need this?
|
||||
self.env.reset_idx(torch.arange(self.num_envs, device=self.env.device))
|
||||
self.env.cumulatives["rewards"][:] = 0
|
||||
self.env.cumulatives["success"][:] = 0
|
||||
obs_dict = self.env.reset()
|
||||
return obs_dict["obs"].detach()
|
||||
|
||||
def step(
|
||||
self, actions: torch.Tensor
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict]:
|
||||
"""Step the environment."""
|
||||
assert isinstance(actions, torch.Tensor)
|
||||
|
||||
# TODO: Check if we need no_grad and detach here
|
||||
with torch.no_grad():
|
||||
obs_dict, rew, dones, infos = self.env.step(actions.detach())
|
||||
truncations = infos["time_outs"]
|
||||
info_ret = {"time_outs": truncations.detach()}
|
||||
if "episode" in infos:
|
||||
info_ret["episode"] = infos["episode"]
|
||||
# NOTE: There's really no way to get the raw observations from IsaacGym
|
||||
# We just use the 'reset_obs' as next_obs, unfortunately.
|
||||
info_ret["observations"] = {"raw": {"obs": obs_dict["obs"].detach()}}
|
||||
return obs_dict["obs"].detach(), rew.detach(), dones.detach(), info_ret
|
||||
|
||||
def render(self):
|
||||
raise NotImplementedError(
|
||||
"We don't support rendering for IsaacLab environments"
|
||||
)
|
||||
|
||||
|
||||
MTBENCH_MW2_CONFIG = {
|
||||
"name": "meta-world-v2",
|
||||
"physics_engine": "physx",
|
||||
"env": {
|
||||
"numEnvs": 1,
|
||||
"envSpacing": 1.5,
|
||||
"episodeLength": 150,
|
||||
"enableDebugVis": False,
|
||||
"clipObservations": 5.0,
|
||||
"clipActions": 1.0,
|
||||
"aggregateMode": 3,
|
||||
"actionScale": 0.01,
|
||||
"resetNoise": 0.15,
|
||||
"tasks": [0],
|
||||
"taskEnvCount": [4096],
|
||||
"init_at_random_progress": True,
|
||||
"exemptedInitAtRandomProgressTasks": [],
|
||||
"taskEmbedding": True,
|
||||
"taskEmbeddingType": "one_hot",
|
||||
"seed": 42,
|
||||
"cameraRenderingInterval": 5000,
|
||||
"cameraWidth": 1024,
|
||||
"cameraHeight": 1024,
|
||||
"sparse_reward": False,
|
||||
"termination_on_success": False,
|
||||
"reward_scale": 1.0,
|
||||
"fixed": False,
|
||||
"numObservations": None,
|
||||
"numActions": 4,
|
||||
},
|
||||
"enableCameraSensors": False,
|
||||
"sim": {
|
||||
"dt": 0.01667,
|
||||
"substeps": 2,
|
||||
"up_axis": "z",
|
||||
"use_gpu_pipeline": True,
|
||||
"gravity": [0.0, 0.0, -9.81],
|
||||
"physx": {
|
||||
"num_threads": 4,
|
||||
"solver_type": 1,
|
||||
"use_gpu": True,
|
||||
"num_position_iterations": 8,
|
||||
"num_velocity_iterations": 1,
|
||||
"contact_offset": 0.005,
|
||||
"rest_offset": 0.0,
|
||||
"bounce_threshold_velocity": 0.2,
|
||||
"max_depenetration_velocity": 1000.0,
|
||||
"default_buffer_size_multiplier": 10.0,
|
||||
"max_gpu_contact_pairs": 1048576,
|
||||
"num_subscenes": 4,
|
||||
"contact_collection": 0,
|
||||
},
|
||||
},
|
||||
"task": {"randomize": False},
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
import jax
|
||||
from mujoco_playground import registry, wrapper_torch
|
||||
|
||||
jax.config.update("jax_compilation_cache_dir", "/tmp/jax_cache")
|
||||
jax.config.update("jax_persistent_cache_min_entry_size_bytes", -1)
|
||||
jax.config.update("jax_persistent_cache_min_compile_time_secs", 0)
|
||||
|
||||
|
||||
class PlaygroundEvalEnvWrapper:
|
||||
def __init__(self, eval_env, max_episode_steps, env_name, num_eval_envs, seed):
|
||||
"""
|
||||
Wrapper used for evaluation / rendering environments.
|
||||
Note that this is different from training environments that are
|
||||
wrapped with RSLRLBraxWrapper.
|
||||
"""
|
||||
self.env = eval_env
|
||||
self.env_name = env_name
|
||||
self.num_envs = num_eval_envs
|
||||
self.jit_reset = jax.jit(jax.vmap(self.env.reset))
|
||||
self.jit_step = jax.jit(jax.vmap(self.env.step))
|
||||
|
||||
if isinstance(self.env.unwrapped.observation_size, dict):
|
||||
self.asymmetric_obs = True
|
||||
else:
|
||||
self.asymmetric_obs = False
|
||||
|
||||
self.key = jax.random.PRNGKey(seed)
|
||||
self.key_reset = jax.random.split(self.key, num_eval_envs)
|
||||
self.max_episode_steps = max_episode_steps
|
||||
|
||||
def reset(self):
|
||||
self.state = self.jit_reset(self.key_reset)
|
||||
if self.asymmetric_obs:
|
||||
obs = wrapper_torch._jax_to_torch(self.state.obs["state"])
|
||||
else:
|
||||
obs = wrapper_torch._jax_to_torch(self.state.obs)
|
||||
return obs
|
||||
|
||||
def step(self, actions):
|
||||
self.state = self.jit_step(self.state, wrapper_torch._torch_to_jax(actions))
|
||||
if self.asymmetric_obs:
|
||||
next_obs = wrapper_torch._jax_to_torch(self.state.obs["state"])
|
||||
else:
|
||||
next_obs = wrapper_torch._jax_to_torch(self.state.obs)
|
||||
rewards = wrapper_torch._jax_to_torch(self.state.reward)
|
||||
dones = wrapper_torch._jax_to_torch(self.state.done)
|
||||
return next_obs, rewards, dones, dones, None
|
||||
|
||||
|
||||
class RandomizeInitialWrapper(wrapper_torch.RSLRLBraxWrapper):
|
||||
"""
|
||||
Wrapper to randomize the initial state of the environment.
|
||||
This is useful for domain randomization experiments.
|
||||
"""
|
||||
|
||||
def reset(self):
|
||||
print("Resetting environment with randomization")
|
||||
obs = super().reset()
|
||||
self.env_state.info["steps"] = jax.random.randint(
|
||||
self.key, self.env_state.info["steps"].shape, 0, 1000
|
||||
).astype(jax.numpy.float32)
|
||||
print(obs)
|
||||
return obs
|
||||
|
||||
def reset_with_critic_obs(self):
|
||||
print("Resetting environment with randomization and critic obs")
|
||||
obs, critic_obs = super().reset_with_critic_obs()
|
||||
self.env_state.info["steps"] = jax.random.randint(
|
||||
self.key, self.env_state.info["steps"].shape, 0, 1000
|
||||
).astype(jax.numpy.float32)
|
||||
return obs, critic_obs
|
||||
|
||||
def step(self, action):
|
||||
obs, reward, done, info = super().step(action)
|
||||
return obs, reward, done, done, info
|
||||
|
||||
|
||||
def make_env(
|
||||
env_name,
|
||||
seed,
|
||||
num_envs,
|
||||
num_eval_envs,
|
||||
device_rank,
|
||||
use_tuned_reward=False,
|
||||
use_domain_randomization=False,
|
||||
use_push_randomization=False,
|
||||
):
|
||||
# Make training environment
|
||||
train_env_cfg = registry.get_default_config(env_name)
|
||||
is_humanoid_task = env_name in [
|
||||
"G1JoystickRoughTerrain",
|
||||
"G1JoystickFlatTerrain",
|
||||
"T1JoystickRoughTerrain",
|
||||
"T1JoystickFlatTerrain",
|
||||
]
|
||||
|
||||
if use_tuned_reward and is_humanoid_task:
|
||||
# NOTE: Tuned reward for G1. Used for producing Figure 7 in the paper.
|
||||
# Somehow it works reasonably for T1 as well.
|
||||
# However, see `sim2real.md` for sim-to-real RL with Booster T1
|
||||
train_env_cfg.reward_config.scales.energy = -5e-5
|
||||
train_env_cfg.reward_config.scales.action_rate = -1e-1
|
||||
train_env_cfg.reward_config.scales.torques = -1e-3
|
||||
train_env_cfg.reward_config.scales.pose = -1.0
|
||||
train_env_cfg.reward_config.scales.tracking_ang_vel = 1.25
|
||||
train_env_cfg.reward_config.scales.tracking_lin_vel = 1.25
|
||||
train_env_cfg.reward_config.scales.feet_phase = 1.0
|
||||
train_env_cfg.reward_config.scales.ang_vel_xy = -0.3
|
||||
train_env_cfg.reward_config.scales.orientation = -5.0
|
||||
|
||||
if is_humanoid_task and not use_push_randomization:
|
||||
train_env_cfg.push_config.enable = False
|
||||
train_env_cfg.push_config.magnitude_range = [0.0, 0.0]
|
||||
randomizer = (
|
||||
registry.get_domain_randomizer(env_name) if use_domain_randomization else None
|
||||
)
|
||||
raw_env = registry.load(env_name, config=train_env_cfg)
|
||||
train_env = RandomizeInitialWrapper(
|
||||
raw_env,
|
||||
num_envs,
|
||||
seed,
|
||||
train_env_cfg.episode_length,
|
||||
train_env_cfg.action_repeat,
|
||||
randomization_fn=randomizer,
|
||||
device_rank=device_rank,
|
||||
)
|
||||
|
||||
# Make evaluation environment
|
||||
eval_env_cfg = registry.get_default_config(env_name)
|
||||
if is_humanoid_task and not use_push_randomization:
|
||||
eval_env_cfg.push_config.enable = False
|
||||
eval_env_cfg.push_config.magnitude_range = [0.0, 0.0]
|
||||
eval_env = registry.load(env_name, config=eval_env_cfg)
|
||||
eval_env = PlaygroundEvalEnvWrapper(
|
||||
eval_env, eval_env_cfg.episode_length, env_name, num_eval_envs, seed
|
||||
)
|
||||
|
||||
return train_env, eval_env
|
||||
Reference in New Issue
Block a user