Fixes build errors due to name conflicts

This commit is contained in:
cvoelcker
2025-07-21 18:31:20 -04:00
parent 094ee0c5ba
commit e2f99648ae
26 changed files with 52 additions and 79 deletions
View File
+375
View File
@@ -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