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
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
+3
View File
@@ -0,0 +1,3 @@
import jax
jax.config.update("jax_default_matmul_precision", "highest")
+45
View File
@@ -0,0 +1,45 @@
import functools
import flax.struct as struct
import jax
import jax.numpy as jnp
class NormalizationState(struct.PyTreeNode):
mean: struct.PyTreeNode
var: struct.PyTreeNode
count: int
class Normalizer:
@functools.partial(jax.jit, static_argnums=0)
def init(self, tree: struct.PyTreeNode) -> NormalizationState:
return NormalizationState(
mean=jax.tree.map(lambda x: jnp.zeros(x.shape[1:], dtype=x.dtype), tree),
var=jax.tree.map(lambda x: jnp.ones(x.shape[1:], dtype=x.dtype), tree),
count=0,
)
@functools.partial(jax.jit, static_argnums=0)
def update(
self, state: NormalizationState, tree: struct.PyTreeNode
) -> NormalizationState:
var = jax.tree.map(lambda x: jnp.var(x, axis=0), tree)
mean = jax.tree.map(lambda x: jnp.mean(x, axis=0), tree)
batch_size = jax.tree.reduce(lambda x, y: y.shape[0], tree, 0)
delta = mean - state.mean
count = state.count + batch_size
new_mean = state.mean + delta * batch_size / count
m_a = state.var * state.count
m_b = var * batch_size
M2 = m_a + m_b + jnp.square(delta) * state.count * batch_size / count
return state.replace(mean=new_mean, var=M2 / count, count=count)
@functools.partial(jax.jit, static_argnums=0)
def normalize(
self, state: NormalizationState, tree: struct.PyTreeNode
) -> struct.PyTreeNode:
return jax.tree.map(
lambda x, m, v: (x - m) / jnp.sqrt(v + 1e-8), tree, state.mean, state.var
)
+750
View File
@@ -0,0 +1,750 @@
import logging
import math
import time
import typing
from typing import Callable, Optional
import distrax
import hydra
import jax
import optax
import plotly.graph_objs as go
from flax import nnx, struct
from flax.struct import PyTreeNode
from gymnax.environments.environment import Environment, EnvParams, EnvState
from jax import numpy as jnp
from jax.experimental import checkify
from jax.random import PRNGKey
from omegaconf import DictConfig, OmegaConf
import wandb
from reppo_alg.env_utils.jax_wrappers import (
BraxGymnaxWrapper,
ClipAction,
LogWrapper,
MjxGymnaxWrapper,
)
from reppo_alg.jaxrl import utils
from reppo_alg.jaxrl.normalization import NormalizationState, Normalizer
logging.basicConfig(level=logging.INFO)
## INITIALIZE CLASS STRUCTURES (NETWORKS, STATES, ...)
class Policy(typing.Protocol):
def __call__(
self,
key: jax.random.PRNGKey,
obs: PyTreeNode,
state: Optional[PyTreeNode] = None,
) -> tuple[PyTreeNode, PyTreeNode]:
pass
class PPOConfig(struct.PyTreeNode):
lr: float
gamma: float
lmbda: float
clip_ratio: float
value_coef: float
entropy_coef: float
total_time_steps: int
num_steps: int
num_mini_batches: int
num_envs: int
num_epochs: int
max_grad_norm: float | None
normalize_advantages: bool
normalize_env: bool
anneal_lr: bool
num_eval: int = 25
max_episode_steps: int = 1000
class Transition(struct.PyTreeNode):
obs: jax.Array
critic_obs: jax.Array
action: jax.Array
reward: jax.Array
log_prob: jax.Array
value: jax.Array
done: jax.Array
truncated: jax.Array
info: dict[str, jax.Array]
class PPOTrainState(nnx.TrainState):
iteration: int
time_steps: int
last_env_state: EnvState
last_obs: jax.Array
last_critic_obs: jax.Array
normalization_state: NormalizationState | None = None
critic_normalization_state: NormalizationState | None = None
class PPONetworks(nnx.Module):
def __init__(
self,
obs_dim: int,
critic_obs_dim: int,
action_dim: int,
hidden_dim: int = 64,
*,
rngs: nnx.Rngs,
):
def linear_layer(in_features, out_features, scale=jnp.sqrt(2)):
return nnx.Linear(
in_features=in_features,
out_features=out_features,
kernel_init=nnx.initializers.orthogonal(scale=scale),
bias_init=nnx.initializers.zeros_init(),
rngs=rngs,
)
self.actor_module = nnx.Sequential(
linear_layer(obs_dim, hidden_dim),
nnx.tanh,
linear_layer(hidden_dim, hidden_dim),
nnx.tanh,
linear_layer(hidden_dim, action_dim, scale=0.01),
)
self.log_std = nnx.Param(jnp.zeros(action_dim))
self.critic_module = nnx.Sequential(
linear_layer(critic_obs_dim, hidden_dim),
nnx.tanh,
linear_layer(hidden_dim, hidden_dim),
nnx.tanh,
linear_layer(hidden_dim, 1, scale=1.0),
)
def critic(self, obs: jax.Array) -> jax.Array:
return self.critic_module(obs).squeeze()
def actor(self, obs: jax.Array) -> distrax.Distribution:
loc = self.actor_module(obs)
pi = distrax.MultivariateNormalDiag(
loc=loc, scale_diag=jnp.exp(self.log_std.value)
)
return pi
def make_policy(train_state: PPOTrainState) -> Policy:
normalizer = Normalizer()
def policy(
key: PRNGKey, obs: jax.Array, state: struct.PyTreeNode = None
) -> tuple[jax.Array, jax.Array]:
if train_state.normalization_state is not None:
obs = normalizer.normalize(train_state.normalization_state, obs)
model = nnx.merge(train_state.graphdef, train_state.params)
pi = model.actor(obs)
value = model.critic(obs)
action = pi.sample(seed=key)
log_prob = pi.log_prob(action)
return action, dict(log_prob=log_prob, value=value)
return policy
def make_eval_fn(
env: Environment, max_episode_steps: int
) -> Callable[[jax.random.PRNGKey, Policy], dict[str, float]]:
def evaluation_fn(key: jax.random.PRNGKey, policy: Policy):
def step_env(carry, _):
key, env_state, obs = carry
key, act_key, env_key = jax.random.split(key, 3)
action, _ = policy(act_key, obs)
env_key = jax.random.split(env_key, env.num_envs)
obs, _, env_state, reward, done, info = env.step(
env_key, env_state, action.clip(-1.0 + 1e-4, 1.0 - 1e-4)
)
return (key, env_state, obs), info
key, init_key = jax.random.split(key)
init_key = jax.random.split(init_key, env.num_envs)
obs, _, env_state = env.reset(init_key)
_, infos = jax.lax.scan(
f=step_env,
init=(key, env_state, obs),
xs=None,
length=max_episode_steps,
)
return {
"episode_return": infos["returned_episode_returns"].mean(
where=infos["returned_episode"]
),
"episode_return_std": infos["returned_episode_returns"].std(
where=infos["returned_episode"]
),
"episode_length": infos["returned_episode_lengths"].mean(
where=infos["returned_episode"]
),
"episode_length_std": infos["returned_episode_lengths"].std(
where=infos["returned_episode"]
),
"num_episodes": infos["returned_episode"].sum(),
}
return evaluation_fn
def make_init(
cfg: PPOConfig,
env: Environment,
env_params: EnvParams = None,
) -> PPOTrainState:
def init(key: jax.random.PRNGKey) -> PPOTrainState:
# Number of calls to train_step
num_train_steps = cfg.total_time_steps // (cfg.num_steps * cfg.num_envs)
# Number of calls to train_iter, add 1 if not divisible by eval_interval
eval_interval = int(
(cfg.total_time_steps / (cfg.num_steps * cfg.num_envs)) // cfg.num_eval
)
num_iterations = num_train_steps // eval_interval + int(
num_train_steps % eval_interval != 0
)
key, model_key = jax.random.split(key)
# Intialize the model
networks = PPONetworks(
obs_dim=env.observation_space(env_params)[0].shape[0],
critic_obs_dim=env.observation_space(env_params)[1].shape[0],
action_dim=env.action_space(env_params).shape[0],
rngs=nnx.Rngs(model_key),
)
# Set initial learning rate
if not cfg.anneal_lr:
lr = cfg.lr
else:
num_iterations = cfg.total_time_steps // cfg.num_steps // cfg.num_envs
num_updates = num_iterations * cfg.num_epochs * cfg.num_mini_batches
lr = optax.linear_schedule(cfg.lr, 1e-6, num_updates)
# Initialize the optimizer
if cfg.max_grad_norm is not None:
optimizer = optax.chain(
optax.clip_by_global_norm(cfg.max_grad_norm),
optax.adam(lr),
)
else:
optimizer = optax.adam(lr)
# Reset and fully initialize the environment
key, env_key = jax.random.split(key)
env_key = jax.random.split(env_key, cfg.num_envs)
obs, critic_obs, env_state = env.reset(env_key)
# randomize initial time step to prevent all envs stepping in tandem
_env_state = env_state.unwrapped()
key, randomize_steps_key = jax.random.split(key)
_env_state.info["steps"] = jax.random.randint(
randomize_steps_key,
_env_state.info["steps"].shape,
0,
cfg.max_episode_steps,
).astype(jnp.float32)
env_state.set_env_state(_env_state)
if cfg.normalize_env:
normalizer = Normalizer()
norm_state = normalizer.init(obs)
critic_normalizer = Normalizer()
critic_norm_state = critic_normalizer.init(critic_obs)
obs = normalizer.normalize(norm_state, obs)
critic_obs = critic_normalizer.normalize(critic_norm_state, critic_obs)
else:
norm_state = None
critic_norm_state = None
# Initialize the state observations of the environment
return PPOTrainState.create(
iteration=0,
time_steps=0,
graphdef=nnx.graphdef(networks),
params=nnx.state(networks),
tx=optimizer,
last_env_state=env_state,
last_obs=obs,
last_critic_obs=critic_obs,
normalization_state=norm_state,
critic_normalization_state=critic_norm_state,
)
return init
def make_train_fn(
cfg: PPOConfig,
env: Environment,
env_params: EnvParams = None,
log_callback: Callable[[PPOTrainState, dict[str, jax.Array]], None] = None,
num_seeds: int = 1,
):
# Initialize the environment and wrap it to admit vectorized behavior.
env_params = env_params or env.default_params
env = ClipAction(env)
env = LogWrapper(env, cfg.num_envs)
eval_fn = make_eval_fn(env, cfg.max_episode_steps)
normalizer = Normalizer()
eval_interval = int(
(cfg.total_time_steps / (cfg.num_steps * cfg.num_envs)) // cfg.num_eval
)
def collect_rollout(
key: PRNGKey, train_state: PPOTrainState
) -> tuple[Transition, PPOTrainState]:
model = nnx.merge(train_state.graphdef, train_state.params)
# Take a step in the environment
def step_env(carry, _) -> tuple[tuple, Transition]:
key, env_state, train_state, obs, critic_obs = carry
if cfg.normalize_env:
norm_state = normalizer.update(train_state.normalization_state, obs)
obs = normalizer.normalize(norm_state, obs)
train_state = train_state.replace(normalization_state=norm_state)
critic_obs = normalizer.normalize(
train_state.critic_normalization_state, critic_obs
)
# Select action
key, act_key, step_key = jax.random.split(key, 3)
pi = model.actor(obs)
action = pi.sample(seed=act_key)
# Take a step in the environment
step_key = jax.random.split(step_key, cfg.num_envs)
next_obs, next_critic_obs, next_env_state, reward, done, info = env.step(
step_key, env_state, action.clip(-1.0 + 1e-4, 1.0 - 1e-4)
)
# Record the transition
transition = Transition(
obs=obs,
critic_obs=critic_obs,
action=action,
reward=reward,
log_prob=pi.log_prob(action),
value=model.critic(critic_obs),
done=done,
truncated=next_env_state.truncated,
info=info,
)
return (
key,
next_env_state,
train_state,
next_obs,
next_critic_obs,
), transition
# Collect rollout via lax.scan taking steps in the environment
rollout_state, transitions = jax.lax.scan(
f=step_env,
init=(
key,
train_state.last_env_state,
train_state,
train_state.last_obs,
train_state.last_critic_obs,
),
length=cfg.num_steps,
)
# Aggregate the transitions across all the environments to reset for the next iteration
_, last_env_state, train_state, last_obs, last_critic_obs = rollout_state
train_state = train_state.replace(
last_env_state=last_env_state,
last_obs=last_obs,
last_critic_obs=last_critic_obs,
time_steps=train_state.time_steps + cfg.num_steps * cfg.num_envs,
)
return transitions, train_state
def learn_step(
key: PRNGKey, train_state: PPOTrainState, batch: Transition
) -> tuple[PPOTrainState, dict[str, jax.Array]]:
# Compute advantages and target values
model = nnx.merge(train_state.graphdef, train_state.params)
if cfg.normalize_env:
last_critic_obs = normalizer.normalize(
train_state.critic_normalization_state, train_state.last_critic_obs
)
else:
last_critic_obs = train_state.last_critic_obs
last_value = model.critic(last_critic_obs)
def compute_advantage(carry, transition):
gae, next_value = carry
done = transition.done
truncated = transition.truncated
reward = transition.reward
value = transition.value
delta = reward + cfg.gamma * next_value * (1 - done) - value
gae = delta + cfg.gamma * cfg.lmbda * (1 - done) * gae
truncated_gae = reward + cfg.gamma * next_value - value
gae = jnp.where(truncated, truncated_gae, gae)
return (gae, value), gae
# Compute the advantage using GAE
_, advantages = jax.lax.scan(
compute_advantage,
(jnp.zeros_like(last_value), last_value),
batch,
reverse=True,
)
target_values = advantages + batch.value
data = (batch, advantages, target_values)
# Reshape data to (num_steps * num_envs, ...)
data = jax.tree.map(
lambda x: x.reshape(
(math.floor(cfg.num_steps * cfg.num_envs), *x.shape[2:])
),
data,
)
def update(train_state, key) -> tuple[PPOTrainState, dict[str, jax.Array]]:
def minibatch_update(carry, indices):
idx, train_state = carry
# Sample data at indices from the batch
minibatch, advantages, target_values = jax.tree.map(
lambda x: jnp.take(x, indices, axis=0), data
)
if cfg.normalize_advantages:
advantages = (advantages - jnp.mean(advantages)) / (
jnp.std(advantages) + 1e-8
)
# Define the loss function
def loss_fn(params):
model = nnx.merge(train_state.graphdef, params)
pi = model.actor(minibatch.obs)
value = model.critic(minibatch.critic_obs)
log_prob = pi.log_prob(minibatch.action)
value_pred_clipped = minibatch.value + (
value - minibatch.value
).clip(-cfg.clip_ratio, cfg.clip_ratio)
value_error = jnp.square(value - target_values)
value_error_clipped = jnp.square(value_pred_clipped - target_values)
value_loss = 0.5 * jnp.mean(
(1.0 - minibatch.truncated)
* jnp.maximum(value_error, value_error_clipped)
)
ratio = jnp.exp(log_prob - minibatch.log_prob)
checkify.check(
jnp.allclose(ratio, 1.0) | (idx != 1),
debug=True,
msg="Ratio not equal to 1 on first iteration: {r}",
r=ratio,
)
actor_loss1 = ratio * advantages
actor_loss2 = (
jnp.clip(ratio, 1 - cfg.clip_ratio, 1 + cfg.clip_ratio)
* advantages
)
actor_loss = -jnp.mean(
(1.0 - minibatch.truncated)
* jnp.minimum(actor_loss1, actor_loss2)
)
entropy_loss = jnp.mean(pi.entropy())
loss = (
actor_loss
+ cfg.value_coef * value_loss
- cfg.entropy_coef * entropy_loss
)
return loss, dict(
actor_loss=actor_loss,
value_loss=value_loss,
entropy_loss=entropy_loss,
loss=loss,
mean_value=value.mean(),
mean_log_prob=log_prob.mean(),
mean_advantages=advantages.mean(),
mean_action=minibatch.action.mean(),
mean_reward=minibatch.reward.mean(),
)
grad_fn = jax.value_and_grad(loss_fn, has_aux=True)
output, grads = grad_fn(train_state.params)
# Global gradient norm (all parameters combined)
flat_grads, _ = jax.flatten_util.ravel_pytree(grads)
global_grad_norm = jnp.linalg.norm(flat_grads)
metrics = output[1]
metrics["advantages"] = advantages
metrics["global_grad_norm"] = global_grad_norm
train_state = train_state.apply_gradients(grads)
return (idx + 1, train_state), metrics
# Shuffle data and split into mini-batches
key, shuffle_key = jax.random.split(key)
mini_batch_size = (
math.floor(cfg.num_steps * cfg.num_envs) // cfg.num_mini_batches
)
indices = jax.random.permutation(shuffle_key, cfg.num_steps * cfg.num_envs)
minibatch_idxs = jax.tree.map(
lambda x: x.reshape(
(cfg.num_mini_batches, mini_batch_size, *x.shape[1:])
),
indices,
)
# Run model update for each mini-batch
train_state, metrics = jax.lax.scan(
minibatch_update, train_state, minibatch_idxs
)
# Compute mean metrics across mini-batches
metrics = jax.tree.map(lambda x: x.mean(0), metrics)
return train_state, metrics
# Update the model for a number of epochs
key, train_key = jax.random.split(key)
(_, train_state), update_metrics = jax.lax.scan(
f=update,
init=(1, train_state),
xs=jax.random.split(train_key, cfg.num_epochs),
)
# Get metrics from the last epoch
update_metrics = jax.tree.map(lambda x: x[-1], update_metrics)
return train_state, update_metrics
# Define the training loop
def train_fn(key: PRNGKey) -> tuple[PPOTrainState, dict]:
def train_eval_step(key, train_state):
def train_step(
state: PPOTrainState, key: PRNGKey
) -> tuple[PPOTrainState, dict[str, jax.Array]]:
key, rollout_key, learn_key = jax.random.split(key, 3)
# Collect trajectories from `state`
transitions, state = collect_rollout(key=rollout_key, train_state=state)
# Execute an update to the policy with `transitions`
state, update_metrics = learn_step(
key=learn_key, train_state=state, batch=transitions
)
metrics = {**update_metrics, **update_metrics}
state = state.replace(iteration=state.iteration + 1)
return state, metrics
train_key, eval_key = jax.random.split(key)
train_state, train_metrics = jax.lax.scan(
f=train_step,
init=train_state,
xs=jax.random.split(train_key, eval_interval),
)
train_metrics = jax.tree.map(lambda x: x[-1], train_metrics)
policy = make_policy(train_state)
eval_metrics = eval_fn(eval_key, policy)
metrics = {
"time_step": train_state.time_steps,
**utils.prefix_dict("train", train_metrics),
**utils.prefix_dict("eval", eval_metrics),
}
return train_state, metrics
def loop_body(
train_state: PPOTrainState, key: PRNGKey
) -> tuple[PPOTrainState, dict]:
# Map execution of the train+eval step across num_seeds (will be looped using jax.lax.scan)
key, subkey = jax.random.split(key)
train_state, metrics = jax.vmap(train_eval_step)(
jax.random.split(subkey, num_seeds), train_state
)
jax.debug.callback(log_callback, train_state, metrics)
return train_state, metrics
# Initialize the policy, environment and map that across the number of random seeds
num_train_steps = cfg.total_time_steps // (cfg.num_steps * cfg.num_envs)
num_iterations = num_train_steps // eval_interval + int(
num_train_steps % eval_interval != 0
)
key, init_key = jax.random.split(key)
train_state = jax.vmap(make_init(cfg, env, env_params))(
jax.random.split(init_key, num_seeds)
)
keys = jax.random.split(key, num_iterations)
# Run the training and evaluation loop from the initialized training state
state, metrics = jax.lax.scan(f=loop_body, init=train_state, xs=keys)
return state, metrics
return train_fn
def plot_history(history: list[dict[str, jax.Array]]):
steps = jnp.array([m["time_step"][0] for m in history])
eval_return = jnp.array([m["eval/episode_return"].mean() for m in history])
eval_return_std = jnp.array([m["eval/episode_return"].std() for m in history])
fig = go.Figure(
[
go.Scatter(
x=steps,
y=eval_return,
name="Mean Episode Return",
mode="lines",
line=dict(color="blue"),
showlegend=False,
),
go.Scatter(
x=steps,
y=eval_return + eval_return_std,
name="Upper Bound",
mode="lines",
line=dict(width=0),
showlegend=False,
),
go.Scatter(
x=steps,
y=eval_return - eval_return_std,
name="Lower Bound",
mode="lines",
line=dict(width=0),
fill="tonexty",
fillcolor="rgba(50, 127, 168, 0.3)",
showlegend=False,
),
]
)
fig.update_layout(
xaxis=dict(title=dict(text="Environment Steps")),
)
return fig
def run(cfg: DictConfig):
metric_history = []
# Define callback to log metrics during training
def log_callback(state, metrics):
metrics["sys_time"] = time.perf_counter()
if len(metric_history) > 0:
num_env_steps = state.time_steps[0] - metric_history[-1]["time_step"][0]
seconds = metrics["sys_time"] - metric_history[-1]["sys_time"]
sps = num_env_steps / seconds
else:
sps = 0
metric_history.append(metrics)
episode_return = metrics["eval/episode_return"].mean()
# Use pop() with a default value of None in case 'advantages' key doesn't exist
advantages = metrics.pop("train/advantages", None)
logging.info(
f"step={state.time_steps[0]} episode_return={episode_return:.3f}, sps={sps:.2f}"
)
log_data = {
"eval/episode_return": episode_return,
"train/advantages": wandb.Histogram(advantages),
**jax.tree.map(jnp.mean, utils.filter_prefix("train", metrics)),
}
# Push log data to WandB
wandb.log(log_data, step=state.time_steps[0])
logging.info(OmegaConf.to_yaml(cfg))
# Set up the experimental environment
if cfg.env.type == "brax":
env = BraxGymnaxWrapper(
cfg.env.name
) # , episode_length=cfg.env.max_episode_steps
elif cfg.env.type == "mjx":
env = MjxGymnaxWrapper(cfg.env.name, episode_length=cfg.env.max_episode_steps)
else:
raise ValueError(f"Unknown environment type: {cfg.env.type}")
key = jax.random.PRNGKey(cfg.seed)
train_fn = make_train_fn(
cfg=PPOConfig(**cfg.hyperparameters),
env=env,
log_callback=log_callback,
num_seeds=cfg.num_seeds,
)
for i in range(cfg.trials):
# Initialize WandB reporting
key, train_key = jax.random.split(key)
wandb.init(
mode=cfg.wandb.mode,
project=cfg.wandb.project,
entity=cfg.wandb.entity,
tags=[cfg.name, cfg.env.name, cfg.env.type, *cfg.tags],
config=OmegaConf.to_container(cfg),
name=f"ppo-{cfg.name}-{cfg.env.name.lower()}",
save_code=True,
)
start = time.perf_counter()
train_state, metrics = jax.jit(train_fn)(train_key)
jax.block_until_ready(metrics)
duration = time.perf_counter() - start
# Save metrics and finish the run
logging.info(f"Training took {duration:.2f} seconds.")
# jnp.savez("metrics.npz", **metrics) # TODO: fix the directory here to save to a unique output directory
wandb.finish()
def tune(cfg: DictConfig):
def log_callback(state, metrics):
episode_return = metrics["eval/episode_return"].mean()
t = state.time_steps[0]
wandb.log(
{
"episode_return": episode_return,
},
step=t,
)
env = MjxGymnaxWrapper(cfg.env.name, episode_length=cfg.env.max_episode_steps)
def train_agent():
wandb.init(project=cfg.wandb.project)
run_cfg = OmegaConf.to_container(cfg)
for k, v in dict(wandb.config).items():
run_cfg["experiment"]["hyperparameters"][k] = v
ppo_cfg = PPOConfig(**run_cfg["experiment"]["hyperparameters"])
train_fn = make_train_fn(
cfg=ppo_cfg,
env=env,
log_callback=log_callback,
num_seeds=cfg.num_seeds,
)
train_fn = jax.jit(train_fn)
logging.info(f"Running experiment with params: \n {run_cfg}")
key = jax.random.PRNGKey(cfg.seed)
train_state, metrics = train_fn(key)
jax.block_until_ready(metrics)
sweep_id = wandb.sweep(
sweep={
"name": f"{cfg.name}-{cfg.env.name}",
"method": "bayes",
"metric": {"name": "episode_return", "goal": "maximize"},
"parameters": {
"lr": {
"values": [1e-4, 3e-4, 1e-3],
},
"normalize_env": {
"values": [True, False],
},
},
},
project=cfg.wandb.project,
entity=cfg.wandb.entity,
)
wandb.agent(sweep_id, function=train_agent, count=cfg.tune.num_runs)
@hydra.main(version_base=None, config_path="../../config", config_name="ppo")
def main(cfg: DictConfig):
if cfg.tune:
tune(cfg)
else:
run(cfg)
if __name__ == "__main__":
main()
+924
View File
@@ -0,0 +1,924 @@
import logging
import time
import typing
from typing import Callable
import hydra
import jax
import numpy as np
import optax
import optuna
import plotly.graph_objs as go
from flax import nnx, struct
from flax.struct import PyTreeNode
from gymnax.environments.environment import Environment, EnvParams, EnvState
from jax import numpy as jnp
from jax.random import PRNGKey
from omegaconf import DictConfig, OmegaConf
import wandb
from reppo_alg.env_utils.jax_wrappers import (
BraxGymnaxWrapper,
ClipAction,
LogWrapper,
MjxGymnaxWrapper,
NormalizeVec,
)
from reppo_alg.jaxrl import utils, muon
from reppo_alg.network_utils.jax_models import (
CategoricalCriticNetwork,
CriticNetwork,
SACActorNetworks,
)
logging.basicConfig(level=logging.INFO)
class Policy(typing.Protocol):
def __call__(
self,
key: jax.random.PRNGKey,
obs: PyTreeNode,
) -> tuple[PyTreeNode, PyTreeNode]:
pass
class Transition(struct.PyTreeNode):
obs: jax.Array
critic_obs: jax.Array
action: jax.Array
reward: jax.Array
next_emb: jax.Array
value: jax.Array
done: jax.Array
truncated: jax.Array
importance_weight: jax.Array
info: dict[str, jax.Array]
class ReppoConfig(struct.PyTreeNode):
lr: float
gamma: float
total_time_steps: int
num_steps: int
lmbda: float
lmbda_min: float
num_mini_batches: int
num_envs: int
num_epochs: int
max_grad_norm: float | None
normalize_env: bool
polyak: float
exploration_noise_min: float
exploration_noise_max: float
exploration_base_envs: int
ent_start: float
ent_target_mult: float
kl_start: float
eval_interval: int = 10
num_eval: int = 25
max_episode_steps: int = 1000
critic_hidden_dim: int = 512
actor_hidden_dim: int = 512
vmin: int = -100
vmax: int = 100
num_bins: int = 250
hl_gauss: bool = False
kl_bound: float = 1.0
aux_loss_mult: float = 0.0
update_kl_lagrangian: bool = True
update_entropy_lagrangian: bool = True
use_critic_norm: bool = True
num_critic_encoder_layers: int = 1
num_critic_head_layers: int = 1
num_critic_pred_layers: int = 1
use_simplical_embedding: bool = False
use_actor_norm: bool = True
num_actor_layers: int = 2
actor_min_std: float = 0.05
reduce_kl: bool = True
reverse_kl: bool = False
anneal_lr: bool = False
actor_kl_clip_mode: str = "clipped"
class SACTrainState(struct.PyTreeNode):
critic: nnx.TrainState
actor: nnx.TrainState
actor_target: nnx.TrainState
iteration: int
time_steps: int
last_env_state: EnvState
last_obs: jax.Array
last_critic_obs: jax.Array
def make_policy(
train_state: SACTrainState,
) -> Callable[[jax.Array, jax.Array], tuple[jax.Array, dict]]:
def policy(key: PRNGKey, obs: jax.Array) -> tuple[jax.Array, dict]:
actor_model = nnx.merge(train_state.actor.graphdef, train_state.actor.params)
action: jax.Array = actor_model.det_action(obs)
return action, {}
return policy
def make_eval_fn(
env: Environment, max_episode_steps: int, reward_scale: float = 1.0
) -> Callable[[jax.random.PRNGKey, Policy, PyTreeNode | None], dict[str, float]]:
def evaluation_fn(
key: jax.random.PRNGKey, policy: Policy, norm_state: PyTreeNode | None
):
def step_env(carry, _):
key, env_state, obs = carry
key, act_key, env_key = jax.random.split(key, 3)
action, _ = policy(act_key, obs)
step_key = jax.random.split(env_key, env.num_envs)
obs, _, env_state, reward, done, info = env.step(
step_key, env_state, action
)
return (key, env_state, obs), info
key, init_key = jax.random.split(key)
init_key = jax.random.split(init_key, env.num_envs)
obs, _, env_state = env.reset(init_key, norm_state)
# randomize initial steps
key, env_key = jax.random.split(key)
_, infos = jax.lax.scan(
f=step_env,
init=(key, env_state, obs),
xs=None,
length=max_episode_steps,
)
return {
"episode_return": infos["returned_episode_returns"].mean(
where=infos["returned_episode"]
)
* reward_scale,
"episode_return_std": infos["returned_episode_returns"].std(
where=infos["returned_episode"]
),
"episode_length": infos["returned_episode_lengths"].mean(
where=infos["returned_episode"]
),
"episode_length_std": infos["returned_episode_lengths"].std(
where=infos["returned_episode"]
),
"num_episodes": infos["returned_episode"].sum(),
}
return evaluation_fn
def make_init(
cfg: ReppoConfig,
env: Environment,
env_params: EnvParams = None,
) -> Callable[[jax.Array], SACTrainState]:
def init(key: jax.random.PRNGKey) -> SACTrainState:
# Number of calls to train_step
key, model_key = jax.random.split(key)
actor_networks = SACActorNetworks(
obs_dim=env.observation_space(env_params)[0].shape[0],
action_dim=env.action_space(env_params).shape[0],
hidden_dim=cfg.actor_hidden_dim,
ent_start=cfg.ent_start,
kl_start=cfg.kl_start,
use_norm=cfg.use_actor_norm,
layers=cfg.num_actor_layers,
rngs=nnx.Rngs(model_key),
)
actor_target_networks = SACActorNetworks(
obs_dim=env.observation_space(env_params)[0].shape[0],
action_dim=env.action_space(env_params).shape[0],
hidden_dim=cfg.actor_hidden_dim,
ent_start=cfg.ent_start,
kl_start=cfg.kl_start,
use_norm=cfg.use_actor_norm,
layers=cfg.num_actor_layers,
rngs=nnx.Rngs(model_key),
)
if cfg.hl_gauss:
critic_networks: nnx.Module = CategoricalCriticNetwork(
obs_dim=env.observation_space(env_params)[1].shape[0],
action_dim=env.action_space(env_params).shape[0],
hidden_dim=cfg.critic_hidden_dim,
num_bins=cfg.num_bins,
vmin=cfg.vmin,
vmax=cfg.vmax,
use_norm=cfg.use_critic_norm,
encoder_layers=cfg.num_critic_encoder_layers,
use_simplical_embedding=cfg.use_simplical_embedding,
head_layers=cfg.num_critic_head_layers,
pred_layers=cfg.num_critic_pred_layers,
rngs=nnx.Rngs(model_key),
)
else:
critic_networks: nnx.Module = CriticNetwork(
obs_dim=env.observation_space(env_params)[1].shape[0],
action_dim=env.action_space(env_params).shape[0],
hidden_dim=cfg.critic_hidden_dim,
use_norm=cfg.use_critic_norm,
encoder_layers=cfg.num_critic_encoder_layers,
use_simplical_embedding=cfg.use_simplical_embedding,
head_layers=cfg.num_critic_head_layers,
pred_layers=cfg.num_critic_pred_layers,
rngs=nnx.Rngs(model_key),
)
if not cfg.anneal_lr:
lr = cfg.lr
else:
num_iterations = cfg.total_time_steps // cfg.num_steps // cfg.num_envs
num_updates = num_iterations * cfg.num_epochs * cfg.num_mini_batches
lr = optax.linear_schedule(cfg.lr, 0, num_updates)
if cfg.max_grad_norm is not None:
actor_optimizer = optax.chain(
optax.clip_by_global_norm(cfg.max_grad_norm),
muon.muon(lr), # optax.adam(lr) optax.adam(lr)
)
critic_optimizer = optax.chain(
optax.clip_by_global_norm(cfg.max_grad_norm),
muon.muon(lr), # optax.adam(lr) optax.adam(lr)
)
else:
actor_optimizer = muon.muon(lr) # optax.adam(lr)
critic_optimizer = muon.muon(lr) # optax.adam(lr)
actor_trainstate = nnx.TrainState.create(
graphdef=nnx.graphdef(actor_networks),
params=nnx.state(actor_networks),
tx=actor_optimizer,
)
actor_target_trainstate = nnx.TrainState.create(
graphdef=nnx.graphdef(actor_target_networks),
params=nnx.state(actor_target_networks),
tx=optax.set_to_zero(),
)
critic_trainstate = nnx.TrainState.create(
graphdef=nnx.graphdef(critic_networks),
params=nnx.state(critic_networks),
tx=critic_optimizer,
)
key, env_key = jax.random.split(key)
env_key = jax.random.split(env_key, cfg.num_envs)
obs, critic_obs, env_state = env.reset(key=env_key, params=env_params)
# randomize initial time step to prevent all envs stepping in tandem
_env_state = env_state.unwrapped()
key, randomize_steps_key = jax.random.split(key)
_env_state.info["steps"] = jax.random.randint(
randomize_steps_key,
_env_state.info["steps"].shape,
0,
cfg.max_episode_steps,
).astype(jnp.float32)
env_state.set_env_state(_env_state)
return SACTrainState(
actor=actor_trainstate,
actor_target=actor_target_trainstate,
critic=critic_trainstate,
iteration=0,
time_steps=0,
last_env_state=env_state,
last_obs=obs,
last_critic_obs=critic_obs,
)
return init
def make_train_fn(
cfg: ReppoConfig,
env: Environment,
env_params: EnvParams = None,
log_callback: Callable[[SACTrainState, dict[str, jax.Array]], None] | None = None,
num_seeds: int = 1,
reward_scale: float = 1.0,
):
env_params = env_params # or env.default_params
env = LogWrapper(env, cfg.num_envs)
env = ClipAction(env)
# env = VecEnv(env, cfg.num_envs)
if cfg.normalize_env:
env = NormalizeVec(env)
eval_fn = make_eval_fn(env, cfg.max_episode_steps, reward_scale=reward_scale)
action_size_target = (
jnp.prod(jnp.array(env.action_space(env_params).shape)) * cfg.ent_target_mult
)
def collect_rollout(
key: PRNGKey, train_state: SACTrainState
) -> tuple[Transition, SACTrainState]:
actor_model = nnx.merge(train_state.actor.graphdef, train_state.actor.params)
critic_model = nnx.merge(train_state.critic.graphdef, train_state.critic.params)
offset = (
jnp.arange(cfg.num_envs - cfg.exploration_base_envs)[:, None]
* (cfg.exploration_noise_max - cfg.exploration_noise_min)
/ (cfg.num_envs - cfg.exploration_base_envs)
) + cfg.exploration_noise_min
offset = jnp.concatenate(
[
jnp.ones((cfg.exploration_base_envs, 1)) * cfg.exploration_noise_min,
offset,
],
axis=0,
)
def step_env(carry, _) -> tuple[tuple, Transition]:
key, env_state, train_state, obs, critic_obs = carry
key, act_key, step_key = jax.random.split(key, 3)
step_key = jax.random.split(step_key, cfg.num_envs)
# get policy action
og_pi = actor_model.actor(obs)
pi = actor_model.actor(obs, scale=offset)
action = pi.sample(seed=act_key)
next_obs, next_critic_obs, next_env_state, reward, done, info = env.step(
step_key, env_state, action
)
# compute importance weights
action = jnp.clip(action, -0.999, 0.999)
raw_importance_weight = jnp.nan_to_num(
og_pi.log_prob(action).sum(-1) - pi.log_prob(action).sum(-1),
nan=jnp.log(cfg.lmbda_min),
)
importance_weight = jnp.clip(
raw_importance_weight, min=jnp.log(cfg.lmbda_min), max=jnp.log(1.0)
)
# compute next state embedding and value
next_action, log_prob = actor_model.actor(next_obs).sample_and_log_prob(
seed=act_key
)
next_emb, value = critic_model.forward(next_critic_obs, next_action)
reward = (
reward
- cfg.gamma * log_prob.sum(-1).squeeze() * actor_model.temperature()
)
transition = Transition(
obs=obs,
critic_obs=critic_obs,
action=action,
next_emb=next_emb,
reward=reward,
value=value,
done=done,
truncated=next_env_state.truncated,
info=info,
importance_weight=importance_weight,
)
return (
key,
next_env_state,
train_state,
next_obs,
next_critic_obs,
), transition
rollout_state, transitions = jax.lax.scan(
f=step_env,
init=(
key,
train_state.last_env_state,
train_state,
train_state.last_obs,
train_state.last_critic_obs,
),
length=cfg.num_steps,
)
_, last_env_state, train_state, last_obs, last_critic_obs = rollout_state
train_state = train_state.replace(
last_env_state=last_env_state,
last_obs=last_obs,
last_critic_obs=last_critic_obs,
time_steps=train_state.time_steps + cfg.num_steps * cfg.num_envs,
)
return transitions, train_state
def learn_step(
key: PRNGKey, train_state: SACTrainState, batch: Transition
) -> tuple[SACTrainState, dict[str, jax.Array]]:
# compute n-step lambda estimates
def compute_nstep_lambda(carry, transition):
lambda_return, truncated, importance_weight = carry
# combine importance_weights with TD lambda
done = transition.done
reward = transition.reward
value = transition.value
lambda_sum = (
jnp.exp(importance_weight) * cfg.lmbda * lambda_return
+ (1 - jnp.exp(importance_weight) * cfg.lmbda) * value
)
delta = cfg.gamma * jnp.where(truncated, value, (1.0 - done) * lambda_sum)
lambda_return = reward + delta
truncated = transition.truncated
return (
lambda_return,
truncated,
transition.importance_weight,
), lambda_return
_, target_values = jax.lax.scan(
compute_nstep_lambda,
(
batch.value[-1],
jnp.ones_like(batch.truncated[0]),
jnp.zeros_like(batch.importance_weight[0]),
),
batch,
reverse=True,
)
# Reshape data to (num_steps * num_envs, ...)
data = (batch, target_values)
data = jax.tree.map(
lambda x: x.reshape((cfg.num_steps * cfg.num_envs, *x.shape[2:])), data
)
train_state = train_state.replace(
actor_target=train_state.actor_target.replace(
params=train_state.actor.params
),
)
actor_target_model = nnx.merge(
train_state.actor_target.graphdef, train_state.actor_target.params
)
def update(train_state, key) -> tuple[SACTrainState, dict[str, jax.Array]]:
def minibatch_update(carry, indices):
idx, train_state = carry
# Sample data at indices from the batch
minibatch, target_values = jax.tree.map(
lambda x: jnp.take(x, indices, axis=0), data
)
def critic_loss_fn(params):
critic_model = nnx.merge(train_state.critic.graphdef, params)
critic_pred = critic_model.critic_cat(
minibatch.critic_obs, minibatch.action
).squeeze()
if cfg.hl_gauss:
target_cat = jax.vmap(
utils.hl_gauss, in_axes=(0, None, None, None)
)(target_values, cfg.num_bins, cfg.vmin, cfg.vmax)
critic_update_loss = optax.softmax_cross_entropy(
critic_pred, target_cat
)
else:
critic_update_loss = optax.squared_error(
critic_pred,
target_values,
)
# Aux loss
pred, value = critic_model.forward(
minibatch.critic_obs, minibatch.action
)
aux_loss = jnp.mean(
(1 - minibatch.done.reshape(-1, 1))
* (pred - minibatch.next_emb) ** 2,
axis=-1,
)
# compute l2 error for logging
critic_loss = optax.squared_error(
value,
target_values,
)
critic_loss = jnp.mean(critic_loss)
loss = jnp.mean(
(1.0 - minibatch.truncated)
* (critic_update_loss + cfg.aux_loss_mult * aux_loss)
)
return loss, dict(
value_loss=critic_loss,
critic_update_loss=critic_update_loss,
loss=loss,
aux_loss=aux_loss,
q=critic_pred.mean(),
abs_batch_action=jnp.abs(minibatch.action).mean(),
reward_mean=minibatch.reward.mean(),
target_values=target_values.mean(),
)
def actor_loss(params):
critic_target_model = nnx.merge(
train_state.critic.graphdef,
train_state.critic.params,
)
actor_model = nnx.merge(train_state.actor.graphdef, params)
# SAC actor loss
pi = actor_model.actor(minibatch.obs)
pred_action, log_prob = pi.sample_and_log_prob(seed=key)
value = critic_target_model.critic(
minibatch.critic_obs, pred_action
)
log_prob = log_prob.sum(-1)
entropy = -log_prob
# policy KL constraint
if cfg.reverse_kl:
pi_action, pi_act_log_prob = pi.sample_and_log_prob(
sample_shape=(16,), seed=key
)
pi_action = jnp.clip(pi_action, -1 + 1e-4, 1 - 1e-4)
old_pi = actor_target_model.actor(minibatch.obs)
old_pi_act_log_prob = old_pi.log_prob(pi_action).sum(-1).mean(0)
pi_act_log_prob = pi_act_log_prob.sum(-1).mean(0)
kl = pi_act_log_prob - old_pi_act_log_prob
else:
old_pi_action, old_pi_act_log_prob = actor_target_model.actor(
minibatch.obs
).sample_and_log_prob(sample_shape=(16,), seed=key)
old_pi_action = jnp.clip(old_pi_action, -1 + 1e-4, 1 - 1e-4)
old_pi_act_log_prob = old_pi_act_log_prob.sum(-1).mean(0)
pi_act_log_prob = pi.log_prob(old_pi_action).sum(-1).mean(0)
kl = old_pi_act_log_prob - pi_act_log_prob
lagrangian = actor_model.lagrangian()
if cfg.actor_kl_clip_mode == "full":
actor_loss = (
log_prob * jax.lax.stop_gradient(actor_model.temperature())
- value
+ kl * jax.lax.stop_gradient(lagrangian) * cfg.reduce_kl
)
elif cfg.actor_kl_clip_mode == "clipped":
actor_loss = jnp.where(
kl < cfg.kl_bound,
log_prob * jax.lax.stop_gradient(actor_model.temperature())
- value,
kl * jax.lax.stop_gradient(lagrangian) * cfg.reduce_kl,
)
elif cfg.actor_kl_clip_mode == "value":
actor_loss = (
log_prob * jax.lax.stop_gradient(actor_model.temperature())
- value
)
else:
raise ValueError(
f"Unknown actor loss mode: {cfg.actor_kl_clip_mode}"
)
# SAC target entropy loss
target_entropy = action_size_target + entropy
target_entropy_loss = (
actor_model.temperature()
* jax.lax.stop_gradient(target_entropy)
)
# Lagrangian constraint (follows temperature update)
lagrangian_loss = -lagrangian * jax.lax.stop_gradient(
kl - cfg.kl_bound
)
# total loss
loss = jnp.mean(actor_loss)
if cfg.update_entropy_lagrangian:
loss += jnp.mean(target_entropy_loss)
if cfg.update_kl_lagrangian:
loss += jnp.mean(lagrangian_loss)
return loss, dict(
actor_loss=actor_loss,
loss=loss,
temp=actor_model.temperature(),
abs_batch_action=jnp.abs(minibatch.action).mean(),
abs_pred_action=jnp.abs(pred_action).mean(),
reward_mean=minibatch.reward.mean(),
kl=kl.mean(),
lagrangian=lagrangian,
lagrangian_loss=lagrangian_loss,
entropy=entropy,
entropy_loss=target_entropy_loss,
target_values=target_values.mean(),
)
critic_grad_fn = jax.value_and_grad(critic_loss_fn, has_aux=True)
output, grads = critic_grad_fn(train_state.critic.params)
critic_train_state = train_state.critic.apply_gradients(grads)
train_state = train_state.replace(
critic=critic_train_state,
)
critic_metrics = output[1]
actor_grad_fn = jax.value_and_grad(actor_loss, has_aux=True)
output, grads = actor_grad_fn(train_state.actor.params)
actor_train_state = train_state.actor.apply_gradients(grads)
train_state = train_state.replace(
actor=actor_train_state,
)
actor_metrics = output[1]
return (idx + 1, train_state), {
**critic_metrics,
**actor_metrics,
}
# Shuffle data and split into mini-batches
key, shuffle_key = jax.random.split(key)
mini_batch_size = (cfg.num_steps * cfg.num_envs) // cfg.num_mini_batches
indices = jax.random.permutation(shuffle_key, cfg.num_steps * cfg.num_envs)
minibatch_idxs = jax.tree.map(
lambda x: x.reshape(
(cfg.num_mini_batches, mini_batch_size, *x.shape[1:])
),
indices,
)
# Run model update for each mini-batch
train_state, metrics = jax.lax.scan(
minibatch_update, train_state, minibatch_idxs
)
# Compute mean metrics across mini-batches
metrics = jax.tree.map(lambda x: x.mean(0), metrics)
return train_state, metrics
# Update the model for a number of epochs
key, train_key = jax.random.split(key)
(_, train_state), update_metrics = jax.lax.scan(
f=update,
init=(1, train_state),
xs=jax.random.split(train_key, cfg.num_epochs),
)
# Get metrics from the last epoch
update_metrics = jax.tree.map(lambda x: x[-1], update_metrics)
return train_state, update_metrics
def train_fn(key: PRNGKey, cfg: ReppoConfig) -> tuple[SACTrainState, dict]:
def train_eval_step(key, train_state):
def train_step(
state: SACTrainState, key: PRNGKey
) -> tuple[SACTrainState, dict[str, jax.Array]]:
key, rollout_key, learn_key = jax.random.split(key, 3)
transitions, state = collect_rollout(key=rollout_key, train_state=state)
state, update_metrics = learn_step(
key=learn_key, train_state=state, batch=transitions
)
metrics = {**update_metrics, **update_metrics}
state = state.replace(iteration=state.iteration + 1)
return state, metrics
train_key, eval_key = jax.random.split(key)
eval_interval = int(
(cfg.total_time_steps / (cfg.num_steps * cfg.num_envs)) // cfg.num_eval
)
train_state, train_metrics = jax.lax.scan(
f=train_step,
init=train_state,
xs=jax.random.split(train_key, eval_interval),
)
train_metrics = jax.tree.map(lambda x: x[-1], train_metrics)
policy = make_policy(train_state)
if cfg.normalize_env:
norm_state = train_state.last_env_state
else:
norm_state = None
eval_metrics = eval_fn(eval_key, policy, norm_state)
train_returns = {
"train/episode_return": train_state.last_env_state.info[
"returned_episode_returns"
].mean(),
"train/episode_length": train_state.last_env_state.info[
"returned_episode_lengths"
].mean(),
}
metrics = {
"time_step": train_state.time_steps,
**utils.prefix_dict("train", train_metrics),
**utils.prefix_dict("eval", eval_metrics),
**train_returns,
}
return train_state, metrics
def loop_body(
train_state: SACTrainState, key: PRNGKey
) -> tuple[SACTrainState, dict]:
key, subkey = jax.random.split(key)
train_state, metrics = jax.vmap(train_eval_step)(
jax.random.split(subkey, num_seeds), train_state
)
jax.debug.callback(log_callback, train_state, metrics)
return train_state, metrics
eval_interval = int(
(cfg.total_time_steps / (cfg.num_steps * cfg.num_envs)) // cfg.num_eval
)
num_train_steps = cfg.total_time_steps // (cfg.num_steps * cfg.num_envs)
num_iterations = num_train_steps // eval_interval + int(
num_train_steps % eval_interval != 0
)
key, init_key = jax.random.split(key)
train_state = jax.vmap(make_init(cfg, env, env_params))(
jax.random.split(init_key, num_seeds)
)
keys = jax.random.split(key, num_iterations)
state, metrics = jax.lax.scan(f=loop_body, init=train_state, xs=keys)
return state, metrics
return train_fn
def plot_history(history: list[dict[str, jax.Array]]):
steps = jnp.array([m["time_step"][0] for m in history])
eval_return = jnp.array([m["eval/episode_return"].mean() for m in history])
eval_return_std = jnp.array([m["eval/episode_return"].std() for m in history])
fig = go.Figure(
[
go.Scatter(
x=steps,
y=eval_return,
name="Mean Episode Return",
mode="lines",
line=dict(color="blue"),
showlegend=False,
),
go.Scatter(
x=steps,
y=eval_return + eval_return_std,
name="Upper Bound",
mode="lines",
line=dict(width=0),
showlegend=False,
),
go.Scatter(
x=steps,
y=eval_return - eval_return_std,
name="Lower Bound",
mode="lines",
line=dict(width=0),
fill="tonexty",
fillcolor="rgba(50, 127, 168, 0.3)",
showlegend=False,
),
]
)
fig.update_layout(
xaxis=dict(title=dict(text="Environment Steps")),
)
return fig
# type object
def _get_optuna_type(trial: optuna.Trial, name, values: list):
if all(isinstance(v, int) for v in values):
return trial.suggest_int(name, low=min(values), high=max(values))
elif all(isinstance(v, float) for v in values):
return trial.suggest_float(name, low=min(values), high=max(values))
elif all(isinstance(v, str) for v in values):
return trial.suggest_categorical(name, values)
elif all(isinstance(v, bool) for v in values):
return trial.suggest_categorical(name, [True, False])
else:
raise ValueError("Values must be of the same type (int, float, or str).")
def run(cfg: DictConfig, trial: optuna.Trial | None) -> float:
"""
Run a single trial of the SAC training process with hyperparameter tuning.
Args:
cfg (DictConfig): Configuration for the SAC training.
trial (optuna.Trial | None): Optuna trial object for hyperparameter tuning.
Returns:
float: The mean episode return from the trial.
"""
sweep_metrics = []
if trial is not None:
# Set hyperparameters from the trial
for name, values in cfg.trial_spec.items():
if name in cfg.hyperparameters:
sampled_value = _get_optuna_type(trial, name, values)
# TODO: Why the fuck is this happening
if isinstance(sampled_value, np.float64):
sampled_value = float(sampled_value)
cfg.hyperparameters[name] = sampled_value
else:
raise ValueError(f"Hyperparameter {name} not found in config.")
try:
with open("completed_trials.txt", "r") as f:
completed_trials = int(f.read())
except FileNotFoundError:
completed_trials = 0
metric_history = []
def log_callback(state, metrics):
metrics["sys_time"] = time.perf_counter()
if len(metric_history) > 0:
num_env_steps = state.time_steps[0] - metric_history[-1]["time_step"][0]
seconds = metrics["sys_time"] - metric_history[-1]["sys_time"]
sps = num_env_steps / seconds
else:
sps = 0
metric_history.append(metrics)
episode_return = metrics["eval/episode_return"].mean()
eval_length = metrics["eval/episode_length"].mean()
logging.info(
f"step={state.time_steps[0]} episode_return={episode_return:.3f}, episode_length={eval_length:.3f} sps={sps:.2f}"
)
log_data = {
"eval/episode_return": episode_return,
"eval/episode_length": eval_length,
**jax.tree.map(jnp.mean, utils.filter_prefix("train", metrics)),
}
wandb.log(log_data, step=state.time_steps[0])
# Set up the experiment
if cfg.env.type == "brax":
env = BraxGymnaxWrapper(
cfg.env.name,
episode_length=cfg.env.max_episode_steps,
reward_scaling=cfg.env.reward_scaling,
terminate=cfg.env.terminate,
)
elif cfg.env.type == "mjx":
env = MjxGymnaxWrapper(
cfg.env.name,
episode_length=cfg.env.max_episode_steps,
reward_scale=cfg.env.reward_scaling,
push_distractions=cfg.env.get("push_distractions", False),
asymmetric_observation=cfg.env.get("asymmetric_observation", False),
)
else:
raise ValueError(f"Unknown environment type: {cfg.env.type}")
# build algo config with overrides
train_fn = make_train_fn(
cfg=ReppoConfig(**cfg.hyperparameters),
env=env,
log_callback=log_callback,
num_seeds=cfg.num_seeds,
reward_scale=1.0 / cfg.env.reward_scaling,
)
for i in range(completed_trials, cfg.num_trials):
cfg.seed = cfg.seed + i
wandb.init(
mode=cfg.wandb.mode,
project=cfg.wandb.project,
entity=cfg.wandb.entity,
tags=[
cfg.name,
cfg.env.name,
cfg.env.type,
"hp_tune" if trial is not None else "val",
*cfg.tags,
],
config=OmegaConf.to_container(cfg),
name=f"resampling-{cfg.name}-{cfg.env.name.lower()}",
save_code=True,
)
logging.info(OmegaConf.to_yaml(cfg))
key = jax.random.PRNGKey(cfg.seed)
start = time.perf_counter()
_, metrics = jax.jit(train_fn, static_argnums=(1,))(
key, ReppoConfig(**cfg.hyperparameters)
)
jax.block_until_ready(metrics)
duration = time.perf_counter() - start
# Save metrics and finish the run
logging.info(f"Training took {duration:.2f} seconds.")
jnp.savez("metrics.npz", **metrics)
wandb.finish()
sweep_metrics.append(metrics["eval/episode_return"])
with open("completed_trials.txt", "w") as f:
f.write(str(i))
sweep_metrics_array = jnp.array(sweep_metrics)
return (0.1 * sweep_metrics_array.mean() + sweep_metrics_array[:, -1].mean()).item()
@hydra.main(version_base=None, config_path="../../config", config_name="reppo")
def main(cfg: DictConfig):
run(cfg, trial=None)
if __name__ == "__main__":
main()
+136
View File
@@ -0,0 +1,136 @@
import distrax
import flax
import jax
import jax.numpy as jnp
def describe(values: jnp.ndarray, axis: tuple | int = 0) -> dict[str, jnp.ndarray]:
"""Compute basic statistics for a batch of values."""
return {
"mean": jnp.mean(values, axis=axis),
"std": jnp.std(values, axis=axis),
"min": jnp.min(values, axis=axis),
"max": jnp.max(values, axis=axis),
}
def merge_dicts(*prefix_dicts: tuple[str, dict], sep: str = "/") -> dict:
"""Merge metric dictionaries with a prefix for each key."""
return {
f"{prefix if prefix else ''}{sep if prefix else ''}{key}": value
for prefix, metrics in prefix_dicts
for key, value in metrics.items()
}
def prefix_dict(prefix: str, metrics: dict, sep: str = "/") -> dict:
"""Add a prefix to all keys in a dictionary."""
return {f"{prefix}{sep}{key}": value for key, value in metrics.items()}
def postfix_dict(postfix: str, metrics: dict, sep: str = "/") -> dict:
"""Add a postfix to all keys in a dictionary."""
return {f"{key}{sep}{postfix}": value for key, value in metrics.items()}
def filter_prefix(prefix: str, metrics: dict, sep: str = "/") -> dict:
"""Filter keys in a dictionary by a prefix."""
return {
key: value for key, value in metrics.items() if key.startswith(prefix + sep)
}
def hl_gauss(inp, num_bins, vmin, vmax, epsilon=0.0):
"""Converts a batch of scalars to soft two-hot encoded targets for discrete regression."""
x = jnp.clip(inp, vmin, max=vmax).squeeze() / (1 - epsilon)
bin_width = (vmax - vmin) / (num_bins - 1)
sigma_to_final_sigma_ratio = 0.75
support = jnp.linspace(
vmin - bin_width / 2, vmax + bin_width / 2, num_bins + 1, dtype=jnp.float32
)
sigma = bin_width * sigma_to_final_sigma_ratio
cdf_evals = jax.scipy.special.erf((support - x) / (jnp.sqrt(2) * sigma))
z = cdf_evals[-1] - cdf_evals[0]
target_probs = cdf_evals[1:] - cdf_evals[:-1]
target_probs = (target_probs / z).reshape(*inp.shape[:-1], num_bins)
uniform = jnp.ones_like(target_probs) / num_bins
return (1 - epsilon) * target_probs + epsilon * uniform
@flax.struct.dataclass
class MultiSampleLogProb:
policy_action: jax.Array
policy_action_log_prob: jax.Array
action: jax.Array
def fast_multi_log_prob(
key: jax.Array,
loc: jax.Array,
scale: jax.Array,
offset_scale: jax.Array,
) -> MultiSampleLogProb:
"""Computes 3 samples from a tanh squashed function
- transformed loc and log_prob
- sample with base scale
- sample with scaled scale
Args:
key: JAX PRNG key.
loc: Location of the distribution.
scale: Scale parameter of the distribution.
offset_scale: Offset scale for the distribution.
"""
# log det factor
# sample base gaussian noise with log prob
base_noise, base_log_prob = distrax.Normal(
jnp.zeros_like(loc), scale
).sample_and_log_prob(seed=key)
base_log_prob = jnp.sum(base_log_prob, axis=-1)
# sample with base scale
base_sample = loc + base_noise
base_sample_transformed = jnp.tanh(base_sample)
# numerically stable jax tanh det jacobian https://github.com/tensorflow/probability/commit/ef6bb176e0ebd1cf6e25c6b5cecdd2428c22963f#diff-e120f70e92e6741bca649f04fcd907b7
base_log_prob -= jnp.sum(
2.0 * (jnp.log(2.0) - base_sample - jax.nn.softplus(-2.0 * base_sample)),
axis=-1,
)
return MultiSampleLogProb(
policy_action=base_sample_transformed,
policy_action_log_prob=base_log_prob,
action=jnp.tanh(loc + offset_scale * base_noise),
)
def multi_softmax(x, dim=8, get_logits=False):
inp_shape = x.shape
if dim is not None:
x = x.reshape(*x.shape[:-1], -1, dim)
if get_logits:
x = jax.nn.log_softmax(x, axis=-1)
else:
x = jax.nn.softmax(x, axis=-1)
return x.reshape(*inp_shape)
def multi_log_softmax(x, dim=8):
if dim is not None:
x = x.reshape(*x.shape[:-1], -1, dim)
return jax.nn.log_softmax(x).reshape(x.shape)
else:
return jax.nn.log_softmax(x, axis=-1)
def simplical_softmax_cross_entropy(pred, target, dim=8):
"""Computes the cross-entropy loss for simplical softmax."""
shape = pred.shape[-1]
if dim is not None:
pred = pred.reshape(*pred.shape[:-1], -1, dim)
target = target.reshape(*target.shape[:-1], -1, dim)
return jnp.sum(-target * jax.nn.log_softmax(pred, axis=-1), axis=-1).mean() / (
shape / dim
)
View File
+270
View File
@@ -0,0 +1,270 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
class DistributionalQNetwork(nn.Module):
def __init__(
self,
n_obs: int,
n_act: int,
num_atoms: int,
v_min: float,
v_max: float,
hidden_dim: int,
device: torch.device = None,
):
super().__init__()
self.net = nn.Sequential(
nn.Linear(n_obs + n_act, hidden_dim, device=device),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim // 2, device=device),
nn.ReLU(),
nn.Linear(hidden_dim // 2, hidden_dim // 4, device=device),
nn.ReLU(),
nn.Linear(hidden_dim // 4, num_atoms, device=device),
)
self.v_min = v_min
self.v_max = v_max
self.num_atoms = num_atoms
def forward(self, obs: torch.Tensor, actions: torch.Tensor) -> torch.Tensor:
x = torch.cat([obs, actions], 1)
x = self.net(x)
return x
def projection(
self,
obs: torch.Tensor,
actions: torch.Tensor,
rewards: torch.Tensor,
bootstrap: torch.Tensor,
discount: float,
q_support: torch.Tensor,
device: torch.device,
) -> torch.Tensor:
delta_z = (self.v_max - self.v_min) / (self.num_atoms - 1)
batch_size = rewards.shape[0]
target_z = (
rewards.unsqueeze(1)
+ bootstrap.unsqueeze(1) * discount.unsqueeze(1) * q_support
)
target_z = target_z.clamp(self.v_min, self.v_max)
b = (target_z - self.v_min) / delta_z
low = torch.floor(b).long()
u = torch.ceil(b).long()
l_mask = torch.logical_and((u > 0), (low == u))
u_mask = torch.logical_and((low < (self.num_atoms - 1)), (low == u))
low = torch.where(l_mask, low - 1, low)
u = torch.where(u_mask, u + 1, u)
next_dist = F.softmax(self.forward(obs, actions), dim=1)
proj_dist = torch.zeros_like(next_dist)
offset = (
torch.linspace(
0, (batch_size - 1) * self.num_atoms, batch_size, device=device
)
.unsqueeze(1)
.expand(batch_size, self.num_atoms)
.long()
)
proj_dist.view(-1).index_add_(
0, (low + offset).view(-1), (next_dist * (u.float() - b)).view(-1)
)
proj_dist.view(-1).index_add_(
0, (u + offset).view(-1), (next_dist * (b - low.float())).view(-1)
)
return proj_dist
class Critic(nn.Module):
def __init__(
self,
n_obs: int,
n_act: int,
num_atoms: int,
v_min: float,
v_max: float,
hidden_dim: int,
device: torch.device = None,
):
super().__init__()
self.qnet1 = DistributionalQNetwork(
n_obs=n_obs,
n_act=n_act,
num_atoms=num_atoms,
v_min=v_min,
v_max=v_max,
hidden_dim=hidden_dim,
device=device,
)
self.qnet2 = DistributionalQNetwork(
n_obs=n_obs,
n_act=n_act,
num_atoms=num_atoms,
v_min=v_min,
v_max=v_max,
hidden_dim=hidden_dim,
device=device,
)
self.register_buffer(
"q_support", torch.linspace(v_min, v_max, num_atoms, device=device)
)
self.device = device
def forward(self, obs: torch.Tensor, actions: torch.Tensor) -> torch.Tensor:
return self.qnet1(obs, actions), self.qnet2(obs, actions)
def projection(
self,
obs: torch.Tensor,
actions: torch.Tensor,
rewards: torch.Tensor,
bootstrap: torch.Tensor,
discount: float,
) -> torch.Tensor:
"""Projection operation that includes q_support directly"""
q1_proj = self.qnet1.projection(
obs,
actions,
rewards,
bootstrap,
discount,
self.q_support,
self.q_support.device,
)
q2_proj = self.qnet2.projection(
obs,
actions,
rewards,
bootstrap,
discount,
self.q_support,
self.q_support.device,
)
return q1_proj, q2_proj
def get_value(self, probs: torch.Tensor) -> torch.Tensor:
"""Calculate value from logits using support"""
return torch.sum(probs * self.q_support, dim=1)
class Actor(nn.Module):
def __init__(
self,
n_obs: int,
n_act: int,
num_envs: int,
init_scale: float,
hidden_dim: int,
std_min: float = 0.05,
std_max: float = 0.8,
device: torch.device = None,
):
super().__init__()
self.n_act = n_act
self.net = nn.Sequential(
nn.Linear(n_obs, hidden_dim, device=device),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim // 2, device=device),
nn.ReLU(),
nn.Linear(hidden_dim // 2, hidden_dim // 4, device=device),
nn.ReLU(),
)
self.fc_mu = nn.Sequential(
nn.Linear(hidden_dim // 4, n_act, device=device),
nn.Tanh(),
)
nn.init.normal_(self.fc_mu[0].weight, 0.0, init_scale)
nn.init.constant_(self.fc_mu[0].bias, 0.0)
noise_scales = (
torch.rand(num_envs, 1, device=device) * (std_max - std_min) + std_min
)
self.register_buffer("noise_scales", noise_scales)
self.register_buffer("std_min", torch.as_tensor(std_min, device=device))
self.register_buffer("std_max", torch.as_tensor(std_max, device=device))
self.n_envs = num_envs
self.device = device
def forward(self, obs: torch.Tensor) -> torch.Tensor:
x = obs
x = self.net(x)
action = self.fc_mu(x)
return action
def explore(
self, obs: torch.Tensor, dones: torch.Tensor = None, deterministic: bool = False
) -> torch.Tensor:
# If dones is provided, resample noise for environments that are done
if dones is not None and dones.sum() > 0:
# Generate new noise scales for done environments (one per environment)
new_scales = (
torch.rand(self.n_envs, 1, device=obs.device)
* (self.std_max - self.std_min)
+ self.std_min
)
# Update only the noise scales for environments that are done
dones_view = dones.view(-1, 1) > 0
self.noise_scales = torch.where(dones_view, new_scales, self.noise_scales)
act = self(obs)
if deterministic:
return act
noise = torch.randn_like(act) * self.noise_scales
return act + noise
class MultiTaskActor(Actor):
def __init__(self, num_tasks: int, task_embedding_dim: int, *args, **kwargs):
super().__init__(*args, **kwargs)
self.num_tasks = num_tasks
self.task_embedding_dim = task_embedding_dim
self.task_embedding = nn.Embedding(
num_tasks, task_embedding_dim, max_norm=1.0, device=self.device
)
def forward(self, obs: torch.Tensor) -> torch.Tensor:
task_ids_one_hot = obs[..., -self.num_tasks :]
task_indices = torch.argmax(task_ids_one_hot, dim=1)
task_embeddings = self.task_embedding(task_indices)
obs = torch.cat([obs[..., : -self.num_tasks], task_embeddings], dim=-1)
return super().forward(obs)
class MultiTaskCritic(Critic):
def __init__(self, num_tasks: int, task_embedding_dim: int, *args, **kwargs):
super().__init__(*args, **kwargs)
self.num_tasks = num_tasks
self.task_embedding_dim = task_embedding_dim
self.task_embedding = nn.Embedding(
num_tasks, task_embedding_dim, max_norm=1.0, device=self.device
)
def forward(self, obs: torch.Tensor, actions: torch.Tensor) -> torch.Tensor:
task_ids_one_hot = obs[..., -self.num_tasks :]
task_indices = torch.argmax(task_ids_one_hot, dim=1)
task_embeddings = self.task_embedding(task_indices)
obs = torch.cat([obs[..., : -self.num_tasks], task_embeddings], dim=-1)
return super().forward(obs, actions)
def projection(
self,
obs: torch.Tensor,
actions: torch.Tensor,
rewards: torch.Tensor,
bootstrap: torch.Tensor,
discount: float,
) -> torch.Tensor:
task_ids_one_hot = obs[..., -self.num_tasks :]
task_indices = torch.argmax(task_ids_one_hot, dim=1)
task_embeddings = self.task_embedding(task_indices)
obs = torch.cat([obs[..., : -self.num_tasks], task_embeddings], dim=-1)
return super().projection(obs, actions, rewards, bootstrap, discount)
+424
View File
@@ -0,0 +1,424 @@
import math
from typing import Sequence, Union
import distrax
import jax
import jax.numpy as jnp
from flax import nnx
from reppo_alg.jaxrl import utils
def torch_he_uniform(
in_axis: Union[int, Sequence[int]] = -2,
out_axis: Union[int, Sequence[int]] = -1,
batch_axis: Sequence[int] = (),
dtype=jnp.float_,
):
"TODO: push to jax"
return nnx.initializers.variance_scaling(
0.3333,
"fan_in",
"uniform",
in_axis=in_axis,
out_axis=out_axis,
batch_axis=batch_axis,
dtype=dtype,
)
class UnitBallNorm(nnx.Module):
def __call__(self, x: jax.Array) -> jax.Array:
return x / (jnp.linalg.norm(x, axis=-1, keepdims=True) + 1e-8)
def normed_activation_layer(
rngs, in_features, out_features, use_norm=True, activation=nnx.swish
):
layers = [
nnx.Linear(
in_features=in_features,
out_features=out_features,
kernel_init=torch_he_uniform(),
rngs=rngs,
)
]
if use_norm:
layers.append(nnx.RMSNorm(out_features, rngs=rngs))
if activation is not None:
layers.append(activation)
return nnx.Sequential(*layers)
class Identity(nnx.Module):
def __call__(self, x: jax.Array) -> jax.Array:
return x
class FCNN(nnx.Module):
def __init__(
self,
in_features: int,
out_features: int,
hidden_dim: int = 512,
hidden_activation=nnx.swish,
output_activation=None,
use_norm: bool = True,
use_output_norm: bool = False,
layers: int = 2,
input_activation: bool = False,
*,
rngs: nnx.Rngs,
):
if layers == 1:
self.module = normed_activation_layer(
rngs,
in_features,
out_features,
use_norm=use_output_norm,
activation=output_activation,
)
else:
if input_activation:
input_layer = nnx.Sequential(
# nnx.LayerNorm(in_features, rngs=rngs) if use_norm else Identity(),
hidden_activation,
normed_activation_layer(
rngs,
in_features,
hidden_dim,
use_norm=use_norm,
activation=hidden_activation,
),
)
else:
input_layer = nnx.Sequential(
normed_activation_layer(
rngs,
in_features,
hidden_dim,
use_norm=use_norm,
activation=hidden_activation,
)
)
hidden_layers = [
normed_activation_layer(
rngs,
hidden_dim,
hidden_dim,
use_norm=use_norm,
activation=hidden_activation,
)
for _ in range(layers - 2)
]
output_layer = normed_activation_layer(
rngs,
hidden_dim,
out_features,
use_norm=use_output_norm,
activation=output_activation,
)
self.module = nnx.Sequential(
input_layer,
*hidden_layers,
output_layer,
)
def __call__(self, x: jax.Array) -> jax.Array:
return self.module(x)
class CriticNetwork(nnx.Module):
def __init__(
self,
obs_dim: int,
action_dim: int,
hidden_dim: int = 512,
use_norm: bool = True,
use_encoder_norm: bool = False,
use_simplical_embedding: bool = False,
encoder_layers: int = 1,
head_layers: int = 1,
pred_layers: int = 1,
*,
rngs: nnx.Rngs,
):
self.feature_module = FCNN(
in_features=obs_dim + action_dim,
out_features=hidden_dim,
hidden_dim=hidden_dim,
hidden_activation=nnx.swish,
output_activation=utils.multi_softmax if use_simplical_embedding else None,
use_norm=use_norm,
use_output_norm=use_encoder_norm,
layers=encoder_layers,
rngs=rngs,
)
self.critic_module = FCNN(
in_features=hidden_dim,
out_features=1,
hidden_dim=hidden_dim,
hidden_activation=nnx.swish,
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
layers=head_layers,
rngs=rngs,
)
self.pred_module = FCNN(
in_features=hidden_dim,
out_features=hidden_dim,
hidden_dim=hidden_dim,
hidden_activation=nnx.swish,
output_activation=utils.multi_softmax if use_simplical_embedding else None,
use_norm=use_norm,
use_output_norm=False,
layers=pred_layers,
rngs=rngs,
)
def features(self, obs: jax.Array, action: jax.Array):
state = jnp.concatenate([obs, action], axis=-1)
return self.feature_module(state)
def critic_head(self, features: jax.Array) -> jax.Array:
return self.critic_module(features)
def critic(self, obs: jax.Array, action: jax.Array) -> jax.Array:
features = self.features(obs, action)
return self.critic_head(features)
def critic_cat(self, obs: jax.Array, action: jax.Array) -> jax.Array:
features = self.features(obs, action)
return self.critic_head(features)
def forward(self, obs, action):
features = self.features(obs, action)
value = self.critic_head(features)
return self.pred_module(features), value
class CategoricalCriticNetwork(nnx.Module):
def __init__(
self,
obs_dim: int,
action_dim: int,
hidden_dim: int = 512,
use_norm: bool = True,
use_encoder_norm: bool = False,
use_simplical_embedding: bool = False,
encoder_layers: int = 1,
head_layers: int = 1,
pred_layers: int = 1,
num_bins: int = 51,
vmin: float = -10.0,
vmax: float = 10.0,
*,
rngs: nnx.Rngs,
):
self.num_bins = num_bins
self.vmin = vmin
self.vmax = vmax
self.feature_module = FCNN(
in_features=obs_dim + action_dim,
out_features=hidden_dim,
hidden_dim=hidden_dim,
hidden_activation=nnx.swish,
output_activation=utils.multi_softmax if use_simplical_embedding else None,
use_norm=use_norm,
use_output_norm=use_encoder_norm,
layers=encoder_layers,
rngs=rngs,
)
self.critic_module = FCNN(
in_features=hidden_dim,
out_features=self.num_bins,
hidden_dim=hidden_dim,
hidden_activation=nnx.swish,
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
layers=head_layers,
input_activation=not use_simplical_embedding,
rngs=rngs,
)
self.pred_module = FCNN(
in_features=hidden_dim,
out_features=hidden_dim,
hidden_dim=hidden_dim,
hidden_activation=nnx.swish,
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
layers=pred_layers,
input_activation=not use_simplical_embedding,
rngs=rngs,
)
self.zero_dist = nnx.Param(
utils.hl_gauss(jnp.zeros((1,)), num_bins, vmin, vmax)
)
def features(self, obs: jax.Array, action: jax.Array):
state = jnp.concatenate([obs, action], axis=-1)
return self.feature_module(state)
def critic_head(self, features: jax.Array) -> jax.Array:
cat = self.critic_module(features) # + self.zero_dist.value * 40.0
return cat
def critic_cat(self, obs: jax.Array, action: jax.Array) -> jax.Array:
features = self.features(obs, action)
return self.critic_head(features)
def critic(self, obs: jax.Array, action: jax.Array) -> jax.Array:
value_cat = jax.nn.softmax(self.critic_cat(obs, action), axis=-1)
value = value_cat.dot(
jnp.linspace(self.vmin, self.vmax, self.num_bins, endpoint=True)
)
return value
def forward(self, obs, action):
features = self.features(obs, action)
value_cat = jax.nn.softmax(self.critic_head(features), axis=-1)
value = value_cat.dot(
jnp.linspace(self.vmin, self.vmax, self.num_bins, endpoint=True)
)
return self.pred_module(features), value
def __call__(self, obs: jax.Array, action: jax.Array) -> jax.Array:
features = self.features(obs, action)
value_cat = jax.nn.softmax(self.critic_head(features), axis=-1)
value = value_cat.dot(
jnp.linspace(self.vmin, self.vmax, self.num_bins, endpoint=True)
)
pred = self.pred_module(features)
return value, value_cat, pred
class SACActorNetworks(nnx.Module):
def __init__(
self,
obs_dim: int,
action_dim: int,
hidden_dim: int = 512,
ent_start: float = 0.1,
kl_start: float = 0.1,
use_norm: bool = True,
layers: int = 2,
min_std: float = 0.1,
*,
rngs: nnx.Rngs,
):
self.actor_module = FCNN(
in_features=obs_dim,
out_features=action_dim * 2,
hidden_dim=hidden_dim,
hidden_activation=nnx.swish,
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
layers=layers,
input_activation=False,
rngs=rngs,
)
start_value = math.log(ent_start)
kl_start_value = math.log(kl_start)
self.temperature_log_param = nnx.Param(jnp.ones(1) * start_value)
self.lagrangian_log_param = nnx.Param(jnp.ones(1) * kl_start_value)
self.min_std = min_std
def actor(
self, obs: jax.Array, scale: float | jax.Array = 1.0
) -> distrax.Distribution:
loc = self.actor_module(obs)
loc, log_std = jnp.split(loc, 2, axis=-1)
std = (jnp.exp(log_std) + self.min_std) * scale
pi = distrax.Transformed(distrax.Normal(loc=loc, scale=std), distrax.Tanh())
return pi
def det_action(self, obs: jax.Array) -> jax.Array:
loc = self.actor_module(obs)
loc, _ = jnp.split(loc, 2, axis=-1)
return jnp.tanh(loc)
def temperature(self) -> jax.Array:
return jnp.exp(self.temperature_log_param.value)
def lagrangian(self) -> jax.Array:
return jnp.exp(self.lagrangian_log_param.value)
def __call__(self, obs: jax.Array) -> jax.Array:
loc = self.actor_module(obs)
loc, std = jnp.split(loc, 2, axis=-1)
return jnp.tanh(loc), std, self.temperature(), self.lagrangian()
class TD3ActorNetworks(nnx.Module):
def __init__(
self,
obs_dim: int,
action_dim: int,
hidden_dim: int = 512,
ent_start: float = 0.1,
kl_start: float = 0.1,
use_norm: bool = True,
layers: int = 2,
min_std: float = 0.1,
*,
rngs: nnx.Rngs,
):
self.actor_module = FCNN(
in_features=obs_dim,
out_features=action_dim * 2,
hidden_dim=hidden_dim,
hidden_activation=nnx.swish,
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
layers=layers,
input_activation=False,
rngs=rngs,
)
start_value = math.log(ent_start)
kl_start_value = math.log(kl_start)
self.temperature_log_param = nnx.Param(jnp.ones(1) * start_value)
self.lagrangian_log_param = nnx.Param(jnp.ones(1) * kl_start_value)
self.min_std = min_std
def actor(
self, obs: jax.Array, scale: float | jax.Array = 1.0
) -> distrax.Distribution:
loc = self.actor_module(obs)
loc, log_std = jnp.split(loc, 2, axis=-1)
std = (jnp.exp(log_std) + self.min_std) * scale
pi = distrax.Transformed(distrax.Normal(loc=loc, scale=std), distrax.Tanh())
return pi
def det_action(self, obs: jax.Array) -> jax.Array:
loc = self.actor_module(obs)
loc, _ = jnp.split(loc, 2, axis=-1)
return jnp.tanh(loc)
def temperature(self) -> jax.Array:
return jnp.exp(self.temperature_log_param.value)
def lagrangian(self) -> jax.Array:
return jnp.exp(self.lagrangian_log_param.value)
class TD3DeterministicDist(distrax.Distribution):
def __init__(self, loc: jax.Array, scale: float | jax.Array):
self.loc = loc
self.scale = scale
def sample(self, seed=None):
return self.loc + self.scale * jax.random.normal(seed, self.loc.shape)
def log_prob(self, value: jax.Array) -> jax.Array:
return jnp.zeros_like(value)
def sample_and_log_prob(self, *, seed, sample_shape=...):
sample = self.sample(seed=seed)
log_prob = self.log_prob(sample)
return sample, log_prob
+376
View File
@@ -0,0 +1,376 @@
import torch
from torch import nn
from torch.distributions import constraints
from torch.distributions.transforms import Transform
from torch.distributions.normal import Normal
from reppo_alg.torchrl.reppo import hl_gauss
class TanhTransform(Transform):
r"""
Transform via the mapping :math:`y = \tanh(x)`.
It is equivalent to
.. code-block:: python
ComposeTransform(
[
AffineTransform(0.0, 2.0),
SigmoidTransform(),
AffineTransform(-1.0, 2.0),
]
)
However this might not be numerically stable, thus it is recommended to use `TanhTransform`
instead.
Note that one should use `cache_size=1` when it comes to `NaN/Inf` values.
"""
domain = constraints.real
codomain = constraints.interval(-1.0, 1.0)
bijective = True
sign = +1
log2 = torch.log(torch.tensor(2.0)).to(
"cuda" if torch.cuda.is_available() else "cpu"
)
def __eq__(self, other):
return isinstance(other, TanhTransform)
def _call(self, x):
return x.tanh()
def _inverse(self, y):
# We do not clamp to the boundary here as it may degrade the performance of certain algorithms.
# one should use `cache_size=1` instead
return torch.atanh(y)
def log_abs_det_jacobian(self, x, y):
# We use a formula that is more numerically stable, see details in the following link
# https://github.com/tensorflow/probability/blob/master/tensorflow_probability/python/bijectors/tanh.py#L69-L80
return 2.0 * (self.log2 - x - torch.nn.functional.softplus(-2.0 * x))
def get_activation(name):
if name == "gelu":
return nn.GELU()
elif name == "relu":
return nn.ReLU()
elif name == "swish":
return nn.SiLU()
elif name is None:
return nn.Identity()
else:
raise ValueError(f"Unknown activation: {name}")
def normed_activation_layer(
in_features, out_features, use_norm=True, activation="swish", device=None
):
layers = [nn.Linear(in_features, out_features, device=device)]
if use_norm:
layers.append(nn.RMSNorm([out_features], device=device))
if activation is not None:
layers.append(get_activation(activation))
return nn.Sequential(*layers)
class FCNN(nn.Module):
def __init__(
self,
in_features,
out_features,
hidden_dim=256,
hidden_activation="swish",
output_activation=None,
use_norm=True,
use_output_norm=False,
layers=2,
input_activation=False,
device=None,
):
super().__init__()
net = []
if layers == 1:
net.append(
normed_activation_layer(
in_features,
out_features,
use_norm=use_output_norm,
activation=output_activation,
device=device,
)
)
else:
if input_activation:
net.append(get_activation(hidden_activation))
net.append(
normed_activation_layer(
in_features,
hidden_dim,
use_norm=use_norm,
activation=hidden_activation,
device=device,
)
)
for _ in range(layers - 2):
net.append(
normed_activation_layer(
hidden_dim,
hidden_dim,
use_norm=use_norm,
activation=hidden_activation,
device=device,
)
)
net.append(
normed_activation_layer(
hidden_dim,
out_features,
use_norm=use_output_norm,
activation=output_activation,
device=device,
)
)
self.net = nn.Sequential(*net)
def forward(self, x):
return self.net(x)
class CriticNetwork(nn.Module):
def __init__(
self,
n_obs,
n_act,
hidden_dim=256,
use_norm=True,
use_encoder_norm=False,
encoder_layers=1,
head_layers=1,
pred_layers=1,
device=None,
):
super().__init__()
self.feature_module = FCNN(
in_features=n_obs + n_act,
out_features=hidden_dim,
hidden_dim=hidden_dim,
hidden_activation="swish",
output_activation=None,
use_norm=use_norm,
use_output_norm=use_encoder_norm,
layers=encoder_layers,
device=device,
)
self.critic_module = FCNN(
in_features=hidden_dim,
out_features=1,
hidden_dim=hidden_dim,
hidden_activation="swish",
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
layers=head_layers,
device=device,
)
self.pred_module = FCNN(
in_features=hidden_dim,
out_features=hidden_dim,
hidden_dim=hidden_dim,
hidden_activation="swish",
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
layers=pred_layers,
device=device,
)
def features(self, obs, action):
state = torch.cat([obs, action], dim=-1)
return self.feature_module(state)
def critic_head(self, features):
return self.critic_module(features)
def critic(self, obs, action):
features = self.features(obs, action)
return self.critic_head(features)
def forward(self, obs, action):
features = self.features(obs, action)
return self.pred_module(features)
class Critic(nn.Module):
def __init__(
self,
n_obs,
n_act,
num_atoms: int,
vmin: float,
vmax: float,
hidden_dim=256,
use_norm=True,
use_encoder_norm=False,
encoder_layers=1,
head_layers=1,
pred_layers=1,
device=None,
):
super().__init__()
self.num_atoms = num_atoms
self.vmin = vmin
self.vmax = vmax
self.hidden_dim = hidden_dim
self.feature_module = FCNN(
in_features=n_obs + n_act,
out_features=hidden_dim,
hidden_dim=hidden_dim,
hidden_activation="swish",
output_activation=None,
use_norm=use_norm,
use_output_norm=use_encoder_norm,
layers=encoder_layers,
device=device,
)
self.critic_module = FCNN(
in_features=hidden_dim,
out_features=num_atoms,
hidden_dim=hidden_dim,
hidden_activation="swish",
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
input_activation=True,
layers=head_layers,
device=device,
)
self.pred_module = FCNN(
in_features=hidden_dim,
out_features=hidden_dim,
hidden_dim=hidden_dim,
hidden_activation="swish",
output_activation=None,
use_norm=use_norm,
input_activation=True,
use_output_norm=False,
layers=pred_layers,
device=device,
)
self.values = torch.linspace(
vmin, vmax, num_atoms, device=device, dtype=torch.float32
)
zeros = hl_gauss(
torch.zeros(1, device=device), self.vmin, self.vmax, self.num_atoms
)
zeros.requires_grad = True
self.zero_dist = nn.Parameter(
hl_gauss(
torch.zeros(1, device=device), self.vmin, self.vmax, self.num_atoms
)
)
def forward(self, obs, action):
inp = torch.cat([obs, action], dim=-1)
features = self.feature_module(inp)
next_pred = self.pred_module(features)
logits = self.critic_module(features) + 40.9 * self.zero_dist
value_cats = torch.softmax(logits, dim=-1)
value = value_cats @ self.values
return value, logits, next_pred, features
class Actor(nn.Module):
def __init__(
self,
n_obs,
n_act,
ent_start: float,
kl_start: float,
hidden_dim=256,
use_norm=True,
layers=2,
min_std=0.1,
device=None,
):
super().__init__()
self.model = FCNN(
in_features=n_obs,
out_features=2 * n_act,
hidden_dim=hidden_dim,
hidden_activation="swish",
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
layers=layers,
device=device,
)
self.log_temp = nn.Parameter(
torch.log(torch.tensor(ent_start, device=device, dtype=torch.float32))
)
self.log_lagrange = nn.Parameter(
torch.log(torch.tensor(kl_start, device=device, dtype=torch.float32))
)
self.min_std = min_std
def forward(self, obs: torch.Tensor) -> torch.distributions.Distribution:
x = self.model(obs)
mean, log_std = torch.split(x, x.shape[-1] // 2, dim=-1)
std = torch.exp(log_std) + self.min_std
pi = Normal(mean, std, validate_args=False)
transformed_pi = torch.distributions.TransformedDistribution(
pi, [torch.distributions.TanhTransform()]
)
return (
transformed_pi,
torch.tanh(mean),
torch.exp(self.log_temp),
torch.exp(self.log_lagrange),
)
class StochasticPolicy(nn.Module):
def __init__(self, actor: Actor, normalizer: nn.Module = None, *args, **kwargs):
super().__init__(*args, **kwargs)
self.actor = actor
self.normalizer = normalizer
def forward(self, obs: torch.Tensor) -> torch.distributions.Distribution:
if self.normalizer:
obs = self.normalizer(obs)
return self.actor(obs)
class TD3DeterministicPolicy(nn.Module):
def __init__(
self,
n_obs,
n_act,
hidden_dim=256,
use_norm=True,
layers=2,
device=None,
):
super().__init__()
self.model = FCNN(
in_features=n_obs,
out_features=2 * n_act,
hidden_dim=hidden_dim,
hidden_activation="swish",
output_activation=None,
use_norm=use_norm,
use_output_norm=False,
layers=layers,
device=device,
)
def forward(self, obs: torch.Tensor) -> torch.Tensor:
x = self.model(obs)
mean, _ = torch.split(x, x.shape[-1] // 2, dim=-1)
return torch.tanh(mean)
View File
+95
View File
@@ -0,0 +1,95 @@
import torch
from omegaconf import DictConfig
def make_envs(cfg: DictConfig, device: torch.device, seed: int = None) -> tuple:
if cfg.env.type == "humanoid_bench":
from reppo_alg.env_utils.torch_wrappers.humanoid_bench_env import (
HumanoidBenchEnv,
)
envs = HumanoidBenchEnv(
cfg.env.name, cfg.hyperparameters.num_envs, device=device
)
return envs, envs
elif cfg.env.type == "isaaclab":
from reppo_alg.env_utils.torch_wrappers.isaaclab_env import IsaacLabEnv
envs = IsaacLabEnv(
cfg.env.name,
device.type,
cfg.hyperparameters.num_envs,
cfg=seed,
action_bounds=cfg.env.action_bounds,
)
return envs, envs
elif cfg.env.type == "mjx":
from reppo_alg.env_utils.torch_wrappers.mujoco_playground_env import make_env
# TODO: Check if re-using same envs for eval could reduce memory usage
envs, eval_envs = make_env(
env_name=cfg.env.name,
seed=seed,
num_envs=cfg.hyperparameters.num_envs,
num_eval_envs=cfg.hyperparameters.num_envs,
device_rank=cfg.platform.device_rank,
use_domain_randomization=False,
use_push_randomization=True,
)
return envs, eval_envs
elif cfg.env.type == "maniskill":
import gymnasium as gym
import mani_skill.envs # noqa: F401
from mani_skill.utils import gym_utils
from mani_skill.utils.wrappers.flatten import FlattenActionSpaceWrapper
from mani_skill.vector.wrappers.gymnasium import ManiSkillVectorEnv
from reppo_alg.env_utils.torch_wrappers.maniskill_wrapper import (
ManiSkillWrapper,
)
envs = gym.make(
cfg.env.name,
num_envs=cfg.hyperparameters.num_envs,
reconfiguration_freq=None,
**cfg.env.env_kwargs,
)
eval_envs = gym.make(
cfg.env.name,
num_envs=cfg.hyperparameters.num_envs,
reconfiguration_freq=1,
**cfg.env.env_kwargs,
)
cfg.env.max_episode_steps = gym_utils.find_max_episode_steps_value(envs)
# heuristic for setting gamma
cfg.hyperparameters.gamma = 1.0 - 10.0 / cfg.env.max_episode_steps
if isinstance(envs.action_space, gym.spaces.Dict):
envs = FlattenActionSpaceWrapper(envs)
eval_envs = FlattenActionSpaceWrapper(eval_envs)
envs = ManiSkillVectorEnv(
envs,
cfg.hyperparameters.num_envs,
ignore_terminations=not cfg.env.partial_reset,
record_metrics=True,
)
eval_envs = ManiSkillVectorEnv(
eval_envs,
cfg.hyperparameters.num_envs,
ignore_terminations=True,
record_metrics=True,
)
return ManiSkillWrapper(
envs,
max_episode_steps=cfg.env.max_episode_steps,
partial_reset=cfg.env.partial_reset,
device=device.type,
), ManiSkillWrapper(
eval_envs,
max_episode_steps=cfg.env.max_episode_steps,
partial_reset=cfg.env.partial_reset,
device=device.type,
)
else:
raise ValueError(
f"Unknown environment type: {cfg.env.type}. Supported types are 'humanoid_bench', 'isaaclab', 'maniskill', and 'mjx'."
)
+691
View File
@@ -0,0 +1,691 @@
import os
import sys
os.environ["TORCHDYNAMO_INLINE_INBUILT_NN_MODULES"] = "1"
os.environ["OMP_NUM_THREADS"] = "1"
if sys.platform != "darwin":
os.environ["MUJOCO_GL"] = "egl"
else:
os.environ["MUJOCO_GL"] = "glfw"
os.environ["XLA_PYTHON_CLIENT_PREALLOCATE"] = "false"
os.environ["JAX_DEFAULT_MATMUL_PRECISION"] = "highest"
import math
import random
import time
import numpy as np
import tqdm
import wandb
try:
# Required for avoiding IsaacGym import error
import isaacgym
except ImportError:
pass
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from reppo_alg.torchrl.reppo import (
EmpiricalNormalization,
PerTaskRewardNormalizer,
RewardNormalizer,
SimpleReplayBuffer,
save_params,
)
from hyperparams import get_args
from tensordict import TensorDict
from torch.amp import GradScaler, autocast
torch.set_float32_matmul_precision("high")
def main():
args = get_args()
print(args)
run_name = f"{args.env_name}__{args.exp_name}__{args.seed}"
amp_enabled = args.amp and args.cuda and torch.cuda.is_available()
amp_device_type = (
"cuda"
if args.cuda and torch.cuda.is_available()
else "mps"
if args.cuda and torch.backends.mps.is_available()
else "cpu"
)
amp_dtype = torch.bfloat16 if args.amp_dtype == "bf16" else torch.float16
scaler = GradScaler(enabled=amp_enabled and amp_dtype == torch.float16)
if args.use_wandb:
wandb.init(
project=args.project,
name=run_name,
config=vars(args),
save_code=True,
)
random.seed(args.seed)
np.random.seed(args.seed)
torch.manual_seed(args.seed)
torch.backends.cudnn.deterministic = args.torch_deterministic
if not args.cuda:
device = torch.device("cpu")
else:
if torch.cuda.is_available():
device = torch.device(f"cuda:{args.device_rank}")
elif torch.backends.mps.is_available():
device = torch.device(f"mps:{args.device_rank}")
else:
raise ValueError("No GPU available")
print(f"Using device: {device}")
if args.env_name.startswith("h1hand-") or args.env_name.startswith("h1-"):
from reppo_alg.env_utils.torch_wrappers.humanoid_bench_env import (
HumanoidBenchEnv,
)
env_type = "humanoid_bench"
envs = HumanoidBenchEnv(args.env_name, args.num_envs, device=device)
eval_envs = envs
elif args.env_name.startswith("Isaac-"):
from reppo_alg.env_utils.torch_wrappers.isaaclab_env import IsaacLabEnv
env_type = "isaaclab"
envs = IsaacLabEnv(
args.env_name,
device.type,
args.num_envs,
args.seed,
action_bounds=args.action_bounds,
)
eval_envs = envs
elif args.env_name.startswith("MTBench-"):
from reppo_alg.env_utils.torch_wrappers.mtbench_env import MTBenchEnv
env_name = "-".join(args.env_name.split("-")[1:])
env_type = "mtbench"
envs = MTBenchEnv(env_name, args.device_rank, args.num_envs, args.seed)
eval_envs = envs
else:
from reppo_alg.env_utils.torch_wrappers.mujoco_playground_env import make_env
# TODO: Check if re-using same envs for eval could reduce memory usage
env_type = "mujoco_playground"
envs, eval_envs = make_env(
args.env_name,
args.seed,
args.num_envs,
args.num_eval_envs,
args.device_rank,
use_tuned_reward=args.use_tuned_reward,
use_domain_randomization=args.use_domain_randomization,
use_push_randomization=args.use_push_randomization,
)
n_act = envs.num_actions
n_obs = envs.num_obs if isinstance(envs.num_obs, int) else envs.num_obs[0]
if envs.asymmetric_obs:
n_critic_obs = (
envs.num_privileged_obs
if isinstance(envs.num_privileged_obs, int)
else envs.num_privileged_obs[0]
)
else:
n_critic_obs = n_obs
action_low, action_high = -1.0, 1.0
if args.obs_normalization:
obs_normalizer = EmpiricalNormalization(shape=n_obs, device=device)
critic_obs_normalizer = EmpiricalNormalization(
shape=n_critic_obs, device=device
)
else:
obs_normalizer = nn.Identity()
critic_obs_normalizer = nn.Identity()
if args.reward_normalization:
if env_type in ["mtbench"]:
reward_normalizer = PerTaskRewardNormalizer(
num_tasks=envs.num_tasks,
gamma=args.gamma,
device=device,
g_max=min(abs(args.v_min), abs(args.v_max)),
)
else:
reward_normalizer = RewardNormalizer(
gamma=args.gamma,
device=device,
g_max=min(abs(args.v_min), abs(args.v_max)),
)
else:
reward_normalizer = nn.Identity()
actor_kwargs = {
"n_obs": n_obs,
"n_act": n_act,
"num_envs": args.num_envs,
"device": device,
"init_scale": args.init_scale,
"hidden_dim": args.actor_hidden_dim,
}
critic_kwargs = {
"n_obs": n_critic_obs,
"n_act": n_act,
"num_atoms": args.num_atoms,
"v_min": args.v_min,
"v_max": args.v_max,
"hidden_dim": args.critic_hidden_dim,
"device": device,
}
if env_type == "mtbench":
actor_kwargs["n_obs"] = n_obs - envs.num_tasks + args.task_embedding_dim
critic_kwargs["n_obs"] = n_critic_obs - envs.num_tasks + args.task_embedding_dim
actor_kwargs["num_tasks"] = envs.num_tasks
actor_kwargs["task_embedding_dim"] = args.task_embedding_dim
critic_kwargs["num_tasks"] = envs.num_tasks
critic_kwargs["task_embedding_dim"] = args.task_embedding_dim
if args.agent == "fasttd3":
if env_type in ["mtbench"]:
from reppo_alg.network_utils.fast_td3_nets import (
MultiTaskActor,
MultiTaskCritic,
)
actor_cls = MultiTaskActor
critic_cls = MultiTaskCritic
else:
from reppo_alg.network_utils.fast_td3_nets import Actor, Critic
actor_cls = Actor
critic_cls = Critic
print("Using FastTD3")
elif args.agent == "fasttd3_simbav2":
if env_type in ["mtbench"]:
from reppo_alg.network_utils.fast_td3_nets_simbav2 import (
MultiTaskActor,
MultiTaskCritic,
)
actor_cls = MultiTaskActor
critic_cls = MultiTaskCritic
else:
from reppo_alg.network_utils.fast_td3_nets_simbav2 import Actor, Critic
actor_cls = Actor
critic_cls = Critic
print("Using FastTD3 + SimbaV2")
actor_kwargs.pop("init_scale")
actor_kwargs.update(
{
"scaler_init": math.sqrt(2.0 / args.actor_hidden_dim),
"scaler_scale": math.sqrt(2.0 / args.actor_hidden_dim),
"alpha_init": 1.0 / (args.actor_num_blocks + 1),
"alpha_scale": 1.0 / math.sqrt(args.actor_hidden_dim),
"expansion": 4,
"c_shift": 3.0,
"num_blocks": args.actor_num_blocks,
}
)
critic_kwargs.update(
{
"scaler_init": math.sqrt(2.0 / args.critic_hidden_dim),
"scaler_scale": math.sqrt(2.0 / args.critic_hidden_dim),
"alpha_init": 1.0 / (args.critic_num_blocks + 1),
"alpha_scale": 1.0 / math.sqrt(args.critic_hidden_dim),
"num_blocks": args.critic_num_blocks,
"expansion": 4,
"c_shift": 3.0,
}
)
else:
raise ValueError(f"Agent {args.agent} not supported")
actor = actor_cls(**actor_kwargs)
if env_type in ["mtbench"]:
# Python 3.8 doesn't support 'from_module' in tensordict
policy = actor.explore
else:
from tensordict import from_module
actor_detach = actor_cls(**actor_kwargs)
# Copy params to actor_detach without grad
from_module(actor).data.to_module(actor_detach)
policy = actor_detach.explore
qnet = critic_cls(**critic_kwargs)
qnet_target = critic_cls(**critic_kwargs)
qnet_target.load_state_dict(qnet.state_dict())
q_optimizer = optim.AdamW(
list(qnet.parameters()),
lr=torch.tensor(args.critic_learning_rate, device=device),
weight_decay=args.weight_decay,
)
actor_optimizer = optim.AdamW(
list(actor.parameters()),
lr=torch.tensor(args.actor_learning_rate, device=device),
weight_decay=args.weight_decay,
)
# Add learning rate schedulers
q_scheduler = optim.lr_scheduler.CosineAnnealingLR(
q_optimizer,
T_max=args.total_timesteps,
eta_min=torch.tensor(args.critic_learning_rate_end, device=device),
)
actor_scheduler = optim.lr_scheduler.CosineAnnealingLR(
actor_optimizer,
T_max=args.total_timesteps,
eta_min=torch.tensor(args.actor_learning_rate_end, device=device),
)
rb = SimpleReplayBuffer(
n_env=args.num_envs,
buffer_size=args.buffer_size,
n_obs=n_obs,
n_act=n_act,
n_critic_obs=n_critic_obs,
asymmetric_obs=envs.asymmetric_obs,
playground_mode=env_type == "mujoco_playground",
n_steps=args.num_steps,
gamma=args.gamma,
device=device,
)
policy_noise = args.policy_noise
noise_clip = args.noise_clip
def evaluate():
obs_normalizer.eval()
num_eval_envs = eval_envs.num_envs
episode_returns = torch.zeros(num_eval_envs, device=device)
episode_lengths = torch.zeros(num_eval_envs, device=device)
done_masks = torch.zeros(num_eval_envs, dtype=torch.bool, device=device)
if env_type == "isaaclab":
obs = eval_envs.reset(random_start_init=False)
else:
obs = eval_envs.reset()
# Run for a fixed number of steps
for i in range(eval_envs.max_episode_steps):
with (
torch.no_grad(),
autocast(
device_type=amp_device_type, dtype=amp_dtype, enabled=amp_enabled
),
):
obs = normalize_obs(obs)
actions = actor(obs)
next_obs, rewards, dones, _, infos = eval_envs.step(actions.float())
if env_type == "mtbench":
# We only report success rate in MTBench evaluation
rewards = (
infos["episode"]["success"].float() if "episode" in infos else 0.0
)
episode_returns = torch.where(
~done_masks, episode_returns + rewards, episode_returns
)
episode_lengths = torch.where(
~done_masks, episode_lengths + 1, episode_lengths
)
if env_type == "mtbench" and "episode" in infos:
dones = dones | infos["episode"]["success"]
done_masks = torch.logical_or(done_masks, dones)
if done_masks.all():
break
obs = next_obs
obs_normalizer.train()
return episode_returns.mean().item(), episode_lengths.mean().item()
def update_main(data, logs_dict):
with autocast(
device_type=amp_device_type, dtype=amp_dtype, enabled=amp_enabled
):
observations = data["observations"]
next_observations = data["next"]["observations"]
if envs.asymmetric_obs:
critic_observations = data["critic_observations"]
next_critic_observations = data["next"]["critic_observations"]
else:
critic_observations = observations
next_critic_observations = next_observations
actions = data["actions"]
rewards = data["next"]["rewards"]
dones = data["next"]["dones"].bool()
truncations = data["next"]["truncations"].bool()
if args.disable_bootstrap:
bootstrap = (~dones).float()
else:
bootstrap = (truncations | ~dones).float()
clipped_noise = torch.randn_like(actions)
clipped_noise = clipped_noise.mul(policy_noise).clamp(
-noise_clip, noise_clip
)
next_state_actions = (actor(next_observations) + clipped_noise).clamp(
action_low, action_high
)
discount = args.gamma ** data["next"]["effective_n_steps"]
with torch.no_grad():
qf1_next_target_projected, qf2_next_target_projected = (
qnet_target.projection(
next_critic_observations,
next_state_actions,
rewards,
bootstrap,
discount,
)
)
qf1_next_target_value = qnet_target.get_value(qf1_next_target_projected)
qf2_next_target_value = qnet_target.get_value(qf2_next_target_projected)
if args.use_cdq:
qf_next_target_dist = torch.where(
qf1_next_target_value.unsqueeze(1)
< qf2_next_target_value.unsqueeze(1),
qf1_next_target_projected,
qf2_next_target_projected,
)
qf1_next_target_dist = qf2_next_target_dist = qf_next_target_dist
else:
qf1_next_target_dist, qf2_next_target_dist = (
qf1_next_target_projected,
qf2_next_target_projected,
)
qf1, qf2 = qnet(critic_observations, actions)
qf1_loss = -torch.sum(
qf1_next_target_dist * F.log_softmax(qf1, dim=1), dim=1
).mean()
qf2_loss = -torch.sum(
qf2_next_target_dist * F.log_softmax(qf2, dim=1), dim=1
).mean()
qf_loss = qf1_loss + qf2_loss
q_optimizer.zero_grad(set_to_none=True)
scaler.scale(qf_loss).backward()
scaler.unscale_(q_optimizer)
critic_grad_norm = torch.nn.utils.clip_grad_norm_(
qnet.parameters(),
max_norm=args.max_grad_norm if args.max_grad_norm > 0 else float("inf"),
)
scaler.step(q_optimizer)
scaler.update()
q_scheduler.step()
logs_dict["critic_grad_norm"] = critic_grad_norm.detach()
logs_dict["qf_loss"] = qf_loss.detach()
logs_dict["qf_max"] = qf1_next_target_value.max().detach()
logs_dict["qf_min"] = qf1_next_target_value.min().detach()
return logs_dict
def update_pol(data, logs_dict):
with autocast(
device_type=amp_device_type, dtype=amp_dtype, enabled=amp_enabled
):
critic_observations = (
data["critic_observations"]
if envs.asymmetric_obs
else data["observations"]
)
qf1, qf2 = qnet(critic_observations, actor(data["observations"]))
qf1_value = qnet.get_value(F.softmax(qf1, dim=1))
qf2_value = qnet.get_value(F.softmax(qf2, dim=1))
if args.use_cdq:
qf_value = torch.minimum(qf1_value, qf2_value)
else:
qf_value = (qf1_value + qf2_value) / 2.0
actor_loss = -qf_value.mean()
actor_optimizer.zero_grad(set_to_none=True)
scaler.scale(actor_loss).backward()
scaler.unscale_(actor_optimizer)
actor_grad_norm = torch.nn.utils.clip_grad_norm_(
actor.parameters(),
max_norm=args.max_grad_norm if args.max_grad_norm > 0 else float("inf"),
)
scaler.step(actor_optimizer)
scaler.update()
actor_scheduler.step()
logs_dict["actor_grad_norm"] = actor_grad_norm.detach()
logs_dict["actor_loss"] = actor_loss.detach()
return logs_dict
if args.compile:
mode = None
update_main = torch.compile(update_main, mode=mode)
update_pol = torch.compile(update_pol, mode=mode)
policy = torch.compile(policy, mode=mode)
normalize_obs = torch.compile(obs_normalizer.forward, mode=mode)
normalize_critic_obs = torch.compile(critic_obs_normalizer.forward, mode=mode)
if args.reward_normalization:
update_stats = torch.compile(reward_normalizer.update_stats, mode=mode)
normalize_reward = torch.compile(reward_normalizer.forward, mode=mode)
else:
normalize_obs = obs_normalizer.forward
normalize_critic_obs = critic_obs_normalizer.forward
if args.reward_normalization:
update_stats = reward_normalizer.update_stats
normalize_reward = reward_normalizer.forward
if envs.asymmetric_obs:
obs, critic_obs = envs.reset_with_critic_obs()
critic_obs = torch.as_tensor(critic_obs, device=device, dtype=torch.float)
else:
obs = envs.reset()
if args.checkpoint_path:
# Load checkpoint if specified
torch_checkpoint = torch.load(
f"{args.checkpoint_path}", map_location=device, weights_only=False
)
actor.load_state_dict(torch_checkpoint["actor_state_dict"])
obs_normalizer.load_state_dict(torch_checkpoint["obs_normalizer_state"])
critic_obs_normalizer.load_state_dict(
torch_checkpoint["critic_obs_normalizer_state"]
)
qnet.load_state_dict(torch_checkpoint["qnet_state_dict"])
qnet_target.load_state_dict(torch_checkpoint["qnet_target_state_dict"])
global_step = torch_checkpoint["global_step"]
else:
global_step = 0
dones = None
pbar = tqdm.tqdm(total=args.total_timesteps, initial=global_step)
start_time = None
desc = ""
while global_step < args.total_timesteps:
logs_dict = TensorDict()
if (
start_time is None
and global_step >= args.measure_burnin + args.learning_starts
):
start_time = time.time()
measure_burnin = global_step
with (
torch.no_grad(),
autocast(device_type=amp_device_type, dtype=amp_dtype, enabled=amp_enabled),
):
norm_obs = normalize_obs(obs)
actions = policy(obs=norm_obs, dones=dones)
next_obs, rewards, dones, _, infos = envs.step(actions.float())
print(infos["time_outs"])
truncations = infos["time_outs"]
if args.reward_normalization:
if env_type == "mtbench":
task_ids_one_hot = obs[..., -envs.num_tasks :]
task_indices = torch.argmax(task_ids_one_hot, dim=1)
update_stats(rewards, dones.float(), task_ids=task_indices)
else:
update_stats(rewards, dones.float())
if envs.asymmetric_obs:
next_critic_obs = infos["observations"]["critic"]
# Compute 'true' next_obs and next_critic_obs for saving
true_next_obs = torch.where(
dones[:, None] > 0, infos["observations"]["raw"]["obs"], next_obs
)
if envs.asymmetric_obs:
true_next_critic_obs = torch.where(
dones[:, None] > 0,
infos["observations"]["raw"]["critic_obs"],
next_critic_obs,
)
transition = TensorDict(
{
"observations": obs,
"actions": torch.as_tensor(actions, device=device, dtype=torch.float),
"next": {
"observations": true_next_obs,
"rewards": torch.as_tensor(
rewards, device=device, dtype=torch.float
),
"truncations": truncations.long(),
"dones": dones.long(),
},
},
batch_size=(envs.num_envs,),
device=device,
)
if envs.asymmetric_obs:
transition["critic_observations"] = critic_obs
transition["next"]["critic_observations"] = true_next_critic_obs
obs = next_obs
if envs.asymmetric_obs:
critic_obs = next_critic_obs
rb.extend(transition)
batch_size = args.batch_size // args.num_envs
if global_step > args.learning_starts:
for i in range(args.num_updates):
data = rb.sample(batch_size)
data["observations"] = normalize_obs(data["observations"])
data["next"]["observations"] = normalize_obs(
data["next"]["observations"]
)
raw_rewards = data["next"]["rewards"]
if env_type in ["mtbench"] and args.reward_normalization:
# Multi-task reward normalization
task_ids_one_hot = data["observations"][..., -envs.num_tasks :]
task_indices = torch.argmax(task_ids_one_hot, dim=1)
data["next"]["rewards"] = normalize_reward(
raw_rewards, task_ids=task_indices
)
else:
data["next"]["rewards"] = normalize_reward(raw_rewards)
if envs.asymmetric_obs:
data["critic_observations"] = normalize_critic_obs(
data["critic_observations"]
)
data["next"]["critic_observations"] = normalize_critic_obs(
data["next"]["critic_observations"]
)
logs_dict = update_main(data, logs_dict)
if args.num_updates > 1:
if i % args.policy_frequency == 1:
logs_dict = update_pol(data, logs_dict)
else:
if global_step % args.policy_frequency == 0:
logs_dict = update_pol(data, logs_dict)
for param, target_param in zip(
qnet.parameters(), qnet_target.parameters()
):
target_param.data.copy_(
args.tau * param.data + (1 - args.tau) * target_param.data
)
if global_step % 100 == 0 and start_time is not None:
speed = (global_step - measure_burnin) / (time.time() - start_time)
pbar.set_description(f"{speed: 4.4f} sps, " + desc)
with torch.no_grad():
logs = {
"actor_loss": logs_dict["actor_loss"].mean(),
"qf_loss": logs_dict["qf_loss"].mean(),
"qf_max": logs_dict["qf_max"].mean(),
"qf_min": logs_dict["qf_min"].mean(),
"actor_grad_norm": logs_dict["actor_grad_norm"].mean(),
"critic_grad_norm": logs_dict["critic_grad_norm"].mean(),
"env_rewards": rewards.mean(),
"buffer_rewards": raw_rewards.mean(),
}
if args.eval_interval > 0 and global_step % args.eval_interval == 0:
print(f"Evaluating at global step {global_step}")
eval_avg_return, eval_avg_length = evaluate()
if env_type in ["humanoid_bench", "isaaclab", "mtbench"]:
# NOTE: Hacky way of evaluating performance, but just works
obs = envs.reset()
logs["eval_avg_return"] = eval_avg_return
logs["eval_avg_length"] = eval_avg_length
if args.use_wandb:
wandb.log(
{
"speed": speed,
"frame": global_step * args.num_envs,
"critic_lr": q_scheduler.get_last_lr()[0],
"actor_lr": actor_scheduler.get_last_lr()[0],
**logs,
},
step=global_step,
)
if (
args.save_interval > 0
and global_step > 0
and global_step % args.save_interval == 0
):
print(f"Saving model at global step {global_step}")
save_params(
global_step,
actor,
qnet,
qnet_target,
obs_normalizer,
critic_obs_normalizer,
args,
f"models/{run_name}_{global_step}.pt",
)
global_step += 1
pbar.update(1)
save_params(
global_step,
actor,
qnet,
qnet_target,
obs_normalizer,
critic_obs_normalizer,
args,
f"models/{run_name}_final.pt",
)
if __name__ == "__main__":
main()
+543
View File
@@ -0,0 +1,543 @@
import os
from dataclasses import dataclass
import tyro
@dataclass
class BaseArgs:
# Default hyperparameters -- specifically for HumanoidBench
# See MuJoCoPlaygroundArgs for default hyperparameters for MuJoCo Playground
# See IsaacLabArgs for default hyperparameters for IsaacLab
env_name: str = "HumanoidRun"
"""the id of the environment"""
agent: str = "fasttd3"
"""the agent to use: currently support [fasttd3, fasttd3_simbav2]"""
seed: int = 1
"""seed of the experiment"""
torch_deterministic: bool = True
"""if toggled, `torch.backends.cudnn.deterministic=False`"""
cuda: bool = True
"""if toggled, cuda will be enabled by default"""
device_rank: int = 0
"""the rank of the device"""
exp_name: str = os.path.basename(__file__)[: -len(".py")]
"""the name of this experiment"""
project: str = "FastTD3"
"""the project name"""
use_wandb: bool = True
"""whether to use wandb"""
checkpoint_path: str = None
"""the path to the checkpoint file"""
num_envs: int = 128
"""the number of environments to run in parallel"""
num_eval_envs: int = 128
"""the number of evaluation environments to run in parallel (only valid for MuJoCo Playground)"""
total_timesteps: int = 50000
"""total timesteps of the experiments"""
critic_learning_rate: float = 3e-4
"""the learning rate of the critic"""
actor_learning_rate: float = 3e-4
"""the learning rate for the actor"""
critic_learning_rate_end: float = 3e-4
"""the learning rate of the critic at the end of training"""
actor_learning_rate_end: float = 3e-4
"""the learning rate for the actor at the end of training"""
buffer_size: int = 1024 * 50
"""the replay memory buffer size"""
num_steps: int = 1
"""the number of steps to use for the multi-step return"""
gamma: float = 0.99
"""the discount factor gamma"""
tau: float = 0.1
"""target smoothing coefficient (default: 0.005)"""
batch_size: int = 32768
"""the batch size of sample from the replay memory"""
policy_noise: float = 0.001
"""the scale of policy noise"""
std_min: float = 0.001
"""the minimum scale of noise"""
std_max: float = 0.4
"""the maximum scale of noise"""
learning_starts: int = 10
"""timestep to start learning"""
policy_frequency: int = 2
"""the frequency of training policy (delayed)"""
noise_clip: float = 0.5
"""noise clip parameter of the Target Policy Smoothing Regularization"""
num_updates: int = 2
"""the number of updates to perform per step"""
init_scale: float = 0.01
"""the scale of the initial parameters"""
num_atoms: int = 101
"""the number of atoms"""
v_min: float = -250.0
"""the minimum value of the support"""
v_max: float = 250.0
"""the maximum value of the support"""
critic_hidden_dim: int = 1024
"""the hidden dimension of the critic network"""
actor_hidden_dim: int = 512
"""the hidden dimension of the actor network"""
critic_num_blocks: int = 2
"""(SimbaV2 only) the number of blocks in the critic network"""
actor_num_blocks: int = 1
"""(SimbaV2 only) the number of blocks in the actor network"""
use_cdq: bool = True
"""whether to use Clipped Double Q-learning"""
measure_burnin: int = 3
"""Number of burn-in iterations for speed measure."""
eval_interval: int = 2500
"""the interval to evaluate the model"""
render_interval: int = 500000
"""the interval to render the model"""
compile: bool = True
"""whether to use torch.compile."""
compile_mode: str = "reduce-overhead"
"""the mode of torch.compile."""
obs_normalization: bool = True
"""whether to enable observation normalization"""
reward_normalization: bool = False
"""whether to enable reward normalization"""
use_grad_norm_clipping: bool = False
"""whether to use gradient norm clipping."""
max_grad_norm: float = 0.0
"""the maximum gradient norm"""
amp: bool = True
"""whether to use amp"""
amp_dtype: str = "bf16"
"""the dtype of the amp"""
disable_bootstrap: bool = False
"""Whether to disable bootstrap in the critic learning"""
use_domain_randomization: bool = False
"""(Playground only) whether to use domain randomization"""
use_push_randomization: bool = False
"""(Playground only) whether to use push randomization"""
use_tuned_reward: bool = False
"""(Playground only) Use tuned reward for G1"""
action_bounds: float = 1.0
"""(IsaacLab only) the bounds of the action space (-action_bounds, action_bounds)"""
task_embedding_dim: int = 32
"""the dimension of the task embedding"""
weight_decay: float = 0.1
"""the weight decay of the optimizer"""
save_interval: int = 5000
"""the interval to save the model"""
def get_args():
"""
Parse command-line arguments and return the appropriate Args instance based on env_name.
"""
# First, parse all arguments using the base Args class
base_args = tyro.cli(BaseArgs)
# Map environment names to their specific Args classes
# For tasks not here, default hyperparameters are used
# See below links for available task list
# - HumanoidBench (https://arxiv.org/abs/2403.10506)
# - IsaacLab (https://isaac-sim.github.io/IsaacLab/main/source/overview/environments.html)
# - MuJoCo Playground (https://arxiv.org/abs/2502.08844)
env_to_args_class = {
# HumanoidBench
# NOTE: These tasks are not full list of HumanoidBench tasks
"h1hand-reach-v0": H1HandReachArgs,
"h1hand-balance-simple-v0": H1HandBalanceSimpleArgs,
"h1hand-balance-hard-v0": H1HandBalanceHardArgs,
"h1hand-pole-v0": H1HandPoleArgs,
"h1hand-truck-v0": H1HandTruckArgs,
"h1hand-maze-v0": H1HandMazeArgs,
"h1hand-push-v0": H1HandPushArgs,
"h1hand-basketball-v0": H1HandBasketballArgs,
"h1hand-window-v0": H1HandWindowArgs,
"h1hand-package-v0": H1HandPackageArgs,
# MuJoCo Playground
# NOTE: These tasks are not full list of MuJoCo Playground tasks
"G1JoystickFlatTerrain": G1JoystickFlatTerrainArgs,
"G1JoystickRoughTerrain": G1JoystickRoughTerrainArgs,
"T1JoystickFlatTerrain": T1JoystickFlatTerrainArgs,
"T1JoystickRoughTerrain": T1JoystickRoughTerrainArgs,
"LeapCubeReorient": LeapCubeReorientArgs,
"LeapCubeRotateZAxis": LeapCubeRotateZAxisArgs,
"Go1JoystickFlatTerrain": Go1JoystickFlatTerrainArgs,
"Go1JoystickRoughTerrain": Go1JoystickRoughTerrainArgs,
"Go1Getup": Go1GetupArgs,
"CheetahRun": CheetahRunArgs, # NOTE: Example config for DeepMind Control Suite
# IsaacLab
# NOTE: These tasks are not full list of IsaacLab tasks
"Isaac-Lift-Cube-Franka-v0": IsaacLiftCubeFrankaArgs,
"Isaac-Open-Drawer-Franka-v0": IsaacOpenDrawerFrankaArgs,
"Isaac-Velocity-Flat-H1-v0": IsaacVelocityFlatH1Args,
"Isaac-Velocity-Flat-G1-v0": IsaacVelocityFlatG1Args,
"Isaac-Velocity-Rough-H1-v0": IsaacVelocityRoughH1Args,
"Isaac-Velocity-Rough-G1-v0": IsaacVelocityRoughG1Args,
"Isaac-Repose-Cube-Allegro-Direct-v0": IsaacReposeCubeAllegroDirectArgs,
"Isaac-Repose-Cube-Shadow-Direct-v0": IsaacReposeCubeShadowDirectArgs,
# MTBench
"MTBench-meta-world-v2-mt10": MetaWorldMT10Args,
"MTBench-meta-world-v2-mt50": MetaWorldMT50Args,
}
# If the provided env_name has a specific Args class, use it
if base_args.env_name in env_to_args_class:
specific_args_class = env_to_args_class[base_args.env_name]
# Re-parse with the specific class, maintaining any user overrides
specific_args = tyro.cli(specific_args_class)
return specific_args
if base_args.env_name.startswith("h1hand-") or base_args.env_name.startswith("h1-"):
# HumanoidBench
specific_args = tyro.cli(HumanoidBenchArgs)
elif base_args.env_name.startswith("Isaac-"):
# IsaacLab
specific_args = tyro.cli(IsaacLabArgs)
elif base_args.env_name.startswith("MTBench-"):
# MTBench
specific_args = tyro.cli(MTBenchArgs)
else:
# MuJoCo Playground
specific_args = tyro.cli(MuJoCoPlaygroundArgs)
return specific_args
@dataclass
class HumanoidBenchArgs(BaseArgs):
# See HumanoidBench (https://arxiv.org/abs/2403.10506) for available task list
total_timesteps: int = 100000
@dataclass
class H1HandReachArgs(HumanoidBenchArgs):
env_name: str = "h1hand-reach-v0"
v_min: float = -2000.0
v_max: float = 2000.0
@dataclass
class H1HandBalanceSimpleArgs(HumanoidBenchArgs):
env_name: str = "h1hand-balance-simple-v0"
total_timesteps: int = 200000
@dataclass
class H1HandBalanceHardArgs(HumanoidBenchArgs):
env_name: str = "h1hand-balance-hard-v0"
total_timesteps: int = 1000000
@dataclass
class H1HandPoleArgs(HumanoidBenchArgs):
env_name: str = "h1hand-pole-v0"
total_timesteps: int = 150000
@dataclass
class H1HandTruckArgs(HumanoidBenchArgs):
env_name: str = "h1hand-truck-v0"
total_timesteps: int = 500000
@dataclass
class H1HandMazeArgs(HumanoidBenchArgs):
env_name: str = "h1hand-maze-v0"
v_min: float = -1000.0
v_max: float = 1000.0
@dataclass
class H1HandPushArgs(HumanoidBenchArgs):
env_name: str = "h1hand-push-v0"
v_min: float = -1000.0
v_max: float = 1000.0
total_timesteps: int = 1000000
@dataclass
class H1HandBasketballArgs(HumanoidBenchArgs):
env_name: str = "h1hand-basketball-v0"
v_min: float = -2000.0
v_max: float = 2000.0
total_timesteps: int = 250000
@dataclass
class H1HandWindowArgs(HumanoidBenchArgs):
env_name: str = "h1hand-window-v0"
total_timesteps: int = 250000
@dataclass
class H1HandPackageArgs(HumanoidBenchArgs):
env_name: str = "h1hand-package-v0"
v_min: float = -10000.0
v_max: float = 10000.0
@dataclass
class MuJoCoPlaygroundArgs(BaseArgs):
# Default hyperparameters for many of Playground environments
v_min: float = -150.0
v_max: float = 150.0
buffer_size: int = 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.99
@dataclass
class MTBenchArgs(BaseArgs):
# Default hyperparameters for MTBench
reward_normalization: bool = True
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 2048 # 2K is usually enough for MTBench
num_envs: int = 4096
num_eval_envs: int = 4096
gamma: float = 0.97
num_steps: int = 8
@dataclass
class MetaWorldMT10Args(MTBenchArgs):
# This config achieves 97 ~ 98% success rate within 10k steps (15-20 mins on A100)
env_name: str = "MTBench-meta-world-v2-mt10"
num_envs: int = 4096
num_eval_envs: int = 4096
num_steps: int = 8
gamma: float = 0.97
@dataclass
class MetaWorldMT50Args(MTBenchArgs):
# FastTD3 + SimbaV2 achieves >90% success rate within 20k steps (80 mins on A100)
# Performance further improves with more training steps, slowly.
env_name: str = "MTBench-meta-world-v2-mt50"
num_envs: int = 8192
num_eval_envs: int = 8192
num_steps: int = 8
gamma: float = 0.99
@dataclass
class G1JoystickFlatTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "G1JoystickFlatTerrain"
total_timesteps: int = 100000
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 128 # 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.97
critic_hidden_dim: int = 1024
batch_size: int = 8129
@dataclass
class G1JoystickRoughTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "G1JoystickRoughTerrain"
total_timesteps: int = 100000
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 128 # 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.97
critic_hidden_dim: int = 1024
batch_size: int = 8129
@dataclass
class T1JoystickFlatTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "T1JoystickFlatTerrain"
total_timesteps: int = 100000
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 128 # 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.97
critic_hidden_dim: int = 1024
batch_size: int = 8129
@dataclass
class T1JoystickRoughTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "T1JoystickRoughTerrain"
total_timesteps: int = 100000
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 128 # 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.97
critic_hidden_dim: int = 1024
batch_size: int = 8129
@dataclass
class T1LowDofJoystickFlatTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "T1LowDofJoystickFlatTerrain"
total_timesteps: int = 1000000
@dataclass
class T1LowDofJoystickRoughTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "T1LowDofJoystickRoughTerrain"
total_timesteps: int = 1000000
@dataclass
class CheetahRunArgs(MuJoCoPlaygroundArgs):
# NOTE: This config will work for most DMC tasks, though we haven't tested DMC extensively.
# Future research can consider using LayerNorm as we find it sometimes works better for DMC tasks.
env_name: str = "CheetahRun"
num_steps: int = 3
v_min: float = -500.0
v_max: float = 500.0
std_min: float = 0.1
policy_noise: float = 0.1
@dataclass
class Go1JoystickFlatTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "Go1JoystickFlatTerrain"
total_timesteps: int = 50000
std_min: float = 0.2
std_max: float = 0.8
policy_noise: float = 0.2
num_updates: int = 8
@dataclass
class Go1JoystickRoughTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "Go1JoystickRoughTerrain"
total_timesteps: int = 50000
std_min: float = 0.2
std_max: float = 0.8
policy_noise: float = 0.2
num_updates: int = 8
@dataclass
class Go1GetupArgs(MuJoCoPlaygroundArgs):
env_name: str = "Go1Getup"
total_timesteps: int = 50000
std_min: float = 0.2
std_max: float = 0.8
policy_noise: float = 0.2
num_updates: int = 8
@dataclass
class LeapCubeReorientArgs(MuJoCoPlaygroundArgs):
env_name: str = "LeapCubeReorient"
num_steps: int = 3
gamma: float = 0.99
policy_noise: float = 0.2
v_min: float = -50.0
v_max: float = 50.0
use_cdq: bool = False
@dataclass
class LeapCubeRotateZAxisArgs(MuJoCoPlaygroundArgs):
env_name: str = "LeapCubeRotateZAxis"
num_steps: int = 1
policy_noise: float = 0.2
gamma: float = 0.99
v_min: float = -10.0
v_max: float = 10.0
use_cdq: bool = False
@dataclass
class IsaacLabArgs(BaseArgs):
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 1024 * 10
num_envs: int = 4096
num_eval_envs: int = 4096
action_bounds: float = 1.0
std_max: float = 0.4
num_atoms: int = 251
render_interval: int = 0 # IsaacLab does not support rendering in our codebase
total_timesteps: int = 100000
@dataclass
class IsaacLiftCubeFrankaArgs(IsaacLabArgs):
# Value learning is unstable for Lift Cube task Due to brittle reward shaping
# Therefore, we need to disable bootstrap from 'reset_obs' in IsaacLab
# Higher UTD works better for manipulation tasks
env_name: str = "Isaac-Lift-Cube-Franka-v0"
num_updates: int = 8
v_min: float = -50.0
v_max: float = 50.0
std_max: float = 0.8
num_envs: int = 1024
num_eval_envs: int = 1024
action_bounds: float = 3.0
disable_bootstrap: bool = True
total_timesteps: int = 20000
@dataclass
class IsaacOpenDrawerFrankaArgs(IsaacLabArgs):
# Higher UTD works better for manipulation tasks
env_name: str = "Isaac-Open-Drawer-Franka-v0"
v_min: float = -50.0
v_max: float = 50.0
num_updates: int = 8
action_bounds: float = 3.0
total_timesteps: int = 20000
@dataclass
class IsaacVelocityFlatH1Args(IsaacLabArgs):
env_name: str = "Isaac-Velocity-Flat-H1-v0"
num_steps: int = 8
num_updates: int = 4
total_timesteps: int = 75000
@dataclass
class IsaacVelocityFlatG1Args(IsaacLabArgs):
env_name: str = "Isaac-Velocity-Flat-G1-v0"
num_steps: int = 8
num_updates: int = 4
total_timesteps: int = 50000
@dataclass
class IsaacVelocityRoughH1Args(IsaacLabArgs):
env_name: str = "Isaac-Velocity-Rough-H1-v0"
num_steps: int = 8
num_updates: int = 4
buffer_size: int = 1024 * 5 # To reduce memory usage
total_timesteps: int = 50000
@dataclass
class IsaacVelocityRoughG1Args(IsaacLabArgs):
env_name: str = "Isaac-Velocity-Rough-G1-v0"
num_steps: int = 8
num_updates: int = 4
buffer_size: int = 1024 * 5 # To reduce memory usage
total_timesteps: int = 50000
@dataclass
class IsaacReposeCubeAllegroDirectArgs(IsaacLabArgs):
env_name: str = "Isaac-Repose-Cube-Allegro-Direct-v0"
total_timesteps: int = 100000
v_min: float = -500.0
v_max: float = 500.0
@dataclass
class IsaacReposeCubeShadowDirectArgs(IsaacLabArgs):
env_name: str = "Isaac-Repose-Cube-Shadow-Direct-v0"
total_timesteps: int = 100000
v_min: float = -500.0
v_max: float = 500.0
+735
View File
@@ -0,0 +1,735 @@
from dataclasses import dataclass, replace
import functools
import os
import random
import sys
import copy
import time
import numpy as np
import tqdm
from omegaconf import DictConfig, OmegaConf
import wandb
from reppo_alg.torchrl.reppo import EmpiricalNormalization, hl_gauss
try:
# Required for avoiding IsaacGym import error
import isaacgym
except ImportError:
pass
import hydra
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torchinfo import summary
from tensordict import TensorDict
from torch.amp import GradScaler
from reppo_alg.torchrl.envs import make_envs
from reppo_alg.network_utils.torch_models import Actor, Critic
torch.set_float32_matmul_precision("medium")
os.environ["TORCHDYNAMO_INLINE_INBUILT_NN_MODULES"] = "1"
os.environ["OMP_NUM_THREADS"] = "1"
if sys.platform != "darwin":
os.environ["MUJOCO_GL"] = "egl"
else:
os.environ["MUJOCO_GL"] = "glfw"
os.environ["XLA_PYTHON_CLIENT_PREALLOCATE"] = "false"
os.environ["JAX_DEFAULT_MATMUL_PRECISION"] = "highest"
@dataclass(slots=True)
class TrainState:
device: torch.device
obs: torch.Tensor
critic_obs: torch.Tensor
actor: Actor
old_actor: Actor
critic: Critic
normalizer: EmpiricalNormalization
critic_normalizer: EmpiricalNormalization
actor_optimizer: optim.Optimizer
critic_optimizer: optim.Optimizer
scaler: GradScaler
def compile(self):
self.actor.compile()
self.old_actor.compile()
self.critic.compile()
self.normalizer.compile()
self.critic_normalizer.compile()
def get_autocast_context(cfg: DictConfig):
amp_enabled = (
cfg.platform.amp_enabled and cfg.platform.cuda and torch.cuda.is_available()
)
amp_device = (
"cuda"
if cfg.platform.cuda and torch.cuda.is_available()
else "mps"
if cfg.platform.cuda and torch.backends.mps.is_available()
else "cpu"
)
amp_dtype = torch.bfloat16 if cfg.platform.amp_dtype == "bf16" else torch.float32
return functools.partial(
torch.amp.autocast,
device_type=amp_device,
dtype=amp_dtype,
enabled=amp_enabled,
)
def make_collect_fn(cfg: DictConfig, env):
autocast = get_autocast_context(cfg)
asymmetric_obs = env.asymmetric_obs
def collect_fn(
train_state: TrainState,
) -> tuple[TrainState, TensorDict, list[dict]]:
transitions = []
info_list = []
obs = train_state.obs
critic_obs = train_state.critic_obs
for _ in range(cfg.hyperparameters.num_steps):
with autocast():
norm_obs = train_state.normalizer(obs)
norm_critic_obs = train_state.critic_normalizer(critic_obs)
with torch.inference_mode():
pi, _, _, _ = train_state.actor(norm_obs)
actions = pi.sample()
next_obs, rewards, dones, truncations, infos = env.step(actions)
if asymmetric_obs:
next_critic_obs = infos["observations"]["critic"]
else:
next_critic_obs = next_obs
with (
torch.inference_mode(),
autocast(),
):
if (
cfg.env.get("has_final_obs", False)
and cfg.env.get("partial_reset", False)
and "final_observation" in infos
):
_next_obs = infos["final_observation"]
_next_critic_obs = _next_obs
else:
_next_obs = next_obs
_next_critic_obs = next_critic_obs
norm_next_obs = train_state.normalizer(_next_obs)
next_pi, _, temperature, _ = train_state.actor(norm_next_obs)
next_actions = next_pi.sample()
next_log_probs = next_pi.log_prob(
next_actions.clip(-1 + 1e-6, 1 - 1e-6)
).sum(-1)
norm_next_critic_obs = train_state.critic_normalizer(_next_critic_obs)
next_value, _, _, next_embedding = train_state.critic(
norm_next_critic_obs, next_actions
)
rewards = (
rewards - cfg.hyperparameters.gamma * next_log_probs * temperature
)
transitions.append(
TensorDict(
{
"observations": norm_obs,
"critic_observations": norm_critic_obs,
"actions": actions,
"log_probs": pi.log_prob(actions.clip(-0.999, 0.999)).sum(-1),
"rewards": rewards.unsqueeze(-1),
"next_embeddings": next_embedding,
"next_values": next_value.unsqueeze(-1),
"dones": dones.unsqueeze(-1).float(),
"truncations": truncations.unsqueeze(-1).float(),
},
batch_size=(env.num_envs,),
)
)
info_list.append(infos)
obs = next_obs
critic_obs = next_critic_obs
train_state = replace(train_state, obs=obs, critic_obs=critic_obs)
return (
train_state,
torch.stack(transitions, dim=0),
info_list,
)
return collect_fn
def make_postprocess_fn(cfg: DictConfig, env):
@torch.compiler.disable()
def compute_gve(rewards, dones, truncated, next_values, device: torch.device):
gves = []
last_gve = 0
truncated[-1] = 1.0
for t in reversed(range(cfg.hyperparameters.num_steps)):
lambda_sum = (
cfg.hyperparameters.lmbda * last_gve
+ (1.0 - cfg.hyperparameters.lmbda) * next_values[t]
)
delta = cfg.hyperparameters.gamma * torch.where(
truncated[t].bool(), next_values[t], (1.0 - dones[t]) * lambda_sum
)
last_gve = rewards[t] + delta
gves.insert(0, last_gve)
return gves
def postprocess(train_state: TrainState, transition: TensorDict):
gve = compute_gve(
rewards=transition["rewards"],
dones=transition["dones"],
truncated=transition["truncations"],
next_values=transition["next_values"],
device=train_state.device,
)
# Flatten all time and environment dimensions into a single batch dimension
data = TensorDict(
{
"observations": transition["observations"],
"critic_observations": transition["critic_observations"],
"actions": transition["actions"],
"rewards": transition["rewards"],
"next_embeddings": transition["next_embeddings"],
"next_values": transition["next_values"],
"dones": transition["dones"],
"truncations": transition["truncations"],
"gve": torch.stack(gve),
},
batch_size=(
cfg.hyperparameters.num_steps,
cfg.hyperparameters.num_envs,
),
device=train_state.device,
)
return data.float().flatten(0, 1).detach()
return postprocess
def make_critic_update_fn(cfg: DictConfig, train_state: TrainState):
autocast = get_autocast_context(cfg)
def update(data: TensorDict):
qnet = train_state.critic
q_optimizer = train_state.critic_optimizer
with autocast():
critic_observations = data["critic_observations"]
actions = data["actions"]
targets = data["gve"]
target_embeddings = data["next_embeddings"]
truncations = data["truncations"].squeeze(-1)
if cfg.env.get("partial_reset", False):
truncation_mask = torch.ones_like(
truncations, dtype=torch.bool, device=train_state.device
)
else:
truncation_mask = 1.0 - truncations
qf_target_dist = hl_gauss(
targets,
cfg.hyperparameters.vmin,
cfg.hyperparameters.vmax,
cfg.hyperparameters.num_bins,
)
_, qf1, embedding, _ = qnet(critic_observations, actions)
qf_loss = -(
truncation_mask
* torch.sum(qf_target_dist * F.log_softmax(qf1, dim=-1), dim=-1)
).mean()
embedding_loss = (
truncation_mask
* F.mse_loss(
embedding,
target_embeddings,
reduction="none",
).mean(dim=-1)
).mean()
qf_loss = qf_loss + cfg.hyperparameters.aux_loss_mult * embedding_loss
q_optimizer.zero_grad(set_to_none=True)
train_state.scaler.scale(qf_loss).backward()
train_state.scaler.unscale_(q_optimizer)
critic_grad_norm = torch.nn.utils.clip_grad_norm_(
qnet.parameters(), max_norm=cfg.hyperparameters.max_grad_norm
)
train_state.scaler.step(q_optimizer)
train_state.scaler.update()
logs_dict = {
"critic_grad_norm": critic_grad_norm.detach(),
"qf_loss": qf_loss.detach(),
"qf_max": targets.max().detach(),
"qf_min": targets.min().detach(),
"qf_mean": targets.mean().detach(),
"embedding_loss": embedding_loss.detach(),
}
return logs_dict
return update
def make_actor_update_fn(cfg: DictConfig, train_state: TrainState):
autocast = get_autocast_context(cfg)
def update(data: TensorDict):
actor = train_state.actor
old_actor = train_state.old_actor
qnet = train_state.critic
actor_optimizer = train_state.actor_optimizer
scaler = train_state.scaler
critic_obs = data["critic_observations"]
with autocast():
pi, _, temperature, beta = actor(data["observations"])
actions = pi.rsample()
log_probs = pi.log_prob(actions.clip(-1 + 1e-6, 1 - 1e-6)).sum(-1)
entropy = -log_probs
qf, _, _, _ = qnet(critic_obs, actions)
actor_loss = -qf + temperature.detach() * log_probs
# compute KL
old_pi, _, _, _ = old_actor(data["observations"])
old_pi_actions = old_pi.sample((16,)).clip(-1 + 1e-6, 1 - 1e-6)
old_log_probs = old_pi.log_prob(old_pi_actions).sum(-1).mean(0)
new_pi_log_probs = pi.log_prob(old_pi_actions).sum(-1).mean(0)
kl = old_log_probs - new_pi_log_probs
if cfg.hyperparameters.actor_kl_clip_mode == "clipped":
actor_loss = torch.where(
kl < cfg.hyperparameters.kl_bound,
actor_loss,
kl * beta.detach(),
).mean()
elif cfg.hyperparameters.actor_kl_clip_mode == "full":
actor_loss = actor_loss + kl * beta.detach()
elif cfg.hyperparameters.actor_kl_clip_mode == "value":
actor_loss = actor_loss
else:
raise ValueError(
f"Unknown actor kl clip mode: {cfg.hyperparameters.actor_kl_clip_mode}"
)
# temperature updates
target_entropy = (
actions.shape[-1] * cfg.hyperparameters.ent_target_mult
) # -0.5 * np.prod(envs.action_space.shape)
entropy_loss = (target_entropy + entropy).detach().mean() * temperature
lagrangian_loss = (
-beta * (kl - cfg.hyperparameters.kl_bound).mean().detach()
)
actor_loss = (actor_loss + entropy_loss + lagrangian_loss).mean()
actor_optimizer.zero_grad(set_to_none=True)
scaler.scale(actor_loss).backward()
scaler.unscale_(actor_optimizer)
actor_grad_norm = torch.nn.utils.clip_grad_norm_(
actor.parameters(), max_norm=cfg.hyperparameters.max_grad_norm
)
scaler.step(actor_optimizer)
scaler.update()
logs_dict = {
"actor_grad_norm": actor_grad_norm.detach(),
"actor_loss": actor_loss.detach(),
"kl": kl.detach(),
"entropy": entropy.detach(),
"temperature": temperature.detach(),
"lagrangian": beta.detach(),
"entropy_loss": entropy_loss.detach(),
"lagrangian_loss": lagrangian_loss.detach(),
}
return logs_dict
return update
def make_evaluate_fn(cfg: DictConfig, eval_envs):
autocast = get_autocast_context(cfg)
@torch.inference_mode()
def evaluate(
train_state: TrainState, stochastic_eval: bool = False
) -> tuple[int | float | bool, int | float | bool]:
train_state.normalizer.eval()
num_eval_envs = eval_envs.num_envs
episode_returns = torch.zeros(num_eval_envs, device=train_state.device)
episode_lengths = torch.zeros(num_eval_envs, device=train_state.device)
done_masks = torch.zeros(
num_eval_envs, dtype=torch.bool, device=train_state.device
)
if cfg.env.type == "isaaclab" or cfg.env.asymmetric_observation:
obs, _ = eval_envs.reset(random_start_init=False)
else:
obs = eval_envs.reset()
# Run for a fixed number of steps
for i in range(eval_envs.max_episode_steps):
with autocast():
obs = train_state.normalizer(obs)
action_dist, det_actions, _, _ = train_state.actor(obs)
if stochastic_eval:
actions = action_dist.sample()
else:
actions = det_actions
next_obs, rewards, dones, _, infos = eval_envs.step(actions)
episode_returns = torch.where(
~done_masks, episode_returns + rewards, episode_returns
)
episode_lengths = torch.where(
~done_masks, episode_lengths + 1, episode_lengths
)
done_masks = torch.logical_or(done_masks, dones)
if done_masks.all():
break
obs = next_obs
train_state.normalizer.train()
if cfg.env.type == "maniskill":
# combine log_infos
info = {
"info_return": infos["log_info"]["return"].mean(),
"episode_len": infos["log_info"]["episode_len"].float().mean(),
"success": infos["log_info"]["success"].float().mean(),
"return": episode_returns.mean().item(),
}
else:
info = {}
return episode_returns.mean().item(), episode_lengths.mean().item(), info
return evaluate
def configure_platform(cfg: DictConfig) -> DictConfig:
cfg.platform.amp_enabled = (
cfg.platform.amp_enabled and cfg.platform.cuda and torch.cuda.is_available()
)
cfg.platform.amp_device = (
"cuda"
if cfg.platform.cuda and torch.cuda.is_available()
else "mps"
if cfg.platform.cuda and torch.backends.mps.is_available()
else "cpu"
)
return cfg
@hydra.main(
version_base=None,
config_path="../../config",
config_name="reppo",
)
def main(cfg):
cfg = configure_platform(cfg)
run_name = f"{cfg.env.name}_torch_{cfg.seed}"
scaler = GradScaler(
enabled=cfg.platform.amp_enabled and cfg.platform.amp_dtype == torch.float16
)
num_batches = cfg.hyperparameters.num_mini_batches
batch_size = (
cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps // num_batches
)
wandb.init(
project=cfg.wandb.project,
name=run_name,
config=OmegaConf.to_container(cfg),
save_code=True,
)
random.seed(cfg.seed)
np.random.seed(cfg.seed)
torch.manual_seed(cfg.seed)
torch.backends.cudnn.deterministic = cfg.platform.torch_deterministic
if not cfg.platform.cuda:
device = torch.device("cpu")
else:
if torch.cuda.is_available():
device = torch.device(f"cuda:{cfg.platform.device_rank}")
elif torch.backends.mps.is_available():
device = torch.device(f"mps:{cfg.platform.device_rank}")
else:
raise ValueError("No GPU available")
print(f"Using device: {device}")
envs, eval_envs = make_envs(cfg=cfg, device=device, seed=cfg.seed)
n_act = envs.num_actions
n_obs = envs.num_obs if isinstance(envs.num_obs, int) else envs.num_obs[0]
if envs.asymmetric_obs:
n_critic_obs = (
envs.num_privileged_obs
if isinstance(envs.num_privileged_obs, int)
else envs.num_privileged_obs[0]
)
else:
n_critic_obs = n_obs
if cfg.hyperparameters.normalize_env:
obs_normalizer = EmpiricalNormalization(shape=n_obs, device=device)
critic_obs_normalizer = EmpiricalNormalization(
shape=n_critic_obs, device=device
)
else:
obs_normalizer = nn.Identity()
critic_obs_normalizer = nn.Identity()
actor = Actor(
n_obs=n_obs,
n_act=n_act,
ent_start=cfg.hyperparameters.ent_start,
kl_start=cfg.hyperparameters.kl_start,
hidden_dim=cfg.hyperparameters.actor_hidden_dim,
use_norm=cfg.hyperparameters.use_actor_norm,
layers=cfg.hyperparameters.num_actor_layers,
min_std=cfg.hyperparameters.actor_min_std,
device=device,
)
old_actor = copy.deepcopy(actor)
qnet = Critic(
n_obs=n_critic_obs,
n_act=n_act,
num_atoms=cfg.hyperparameters.num_bins,
vmin=cfg.hyperparameters.vmin,
vmax=cfg.hyperparameters.vmax,
hidden_dim=cfg.hyperparameters.critic_hidden_dim,
use_norm=cfg.hyperparameters.use_critic_norm,
use_encoder_norm=False,
encoder_layers=cfg.hyperparameters.num_critic_encoder_layers,
head_layers=cfg.hyperparameters.num_critic_head_layers,
pred_layers=cfg.hyperparameters.num_critic_pred_layers,
device=device,
)
q_optimizer = optim.AdamW(
list(qnet.parameters()),
lr=torch.tensor(cfg.hyperparameters.lr, device=device),
)
actor_optimizer = optim.AdamW(
list(actor.parameters()),
lr=torch.tensor(cfg.hyperparameters.lr, device=device),
)
if envs.asymmetric_obs:
obs, critic_obs = envs.reset_with_critic_obs()
critic_obs = torch.as_tensor(critic_obs, device=device, dtype=torch.float)
else:
obs = envs.reset()
critic_obs = obs
train_state = TrainState(
obs=obs,
critic_obs=critic_obs,
actor=actor,
old_actor=old_actor,
critic=qnet,
normalizer=obs_normalizer,
critic_normalizer=critic_obs_normalizer,
actor_optimizer=actor_optimizer,
critic_optimizer=q_optimizer,
device=device,
scaler=scaler,
)
print(
summary(
train_state.critic,
input_data=(critic_obs[:1], torch.zeros((1, n_act), device=device)),
depth=10,
)
)
print(summary(train_state.actor, input_data=(obs[:1],), depth=10))
# create functions
collect_fn = make_collect_fn(cfg, envs)
postprocess_fn = make_postprocess_fn(cfg, envs)
update_critic = make_critic_update_fn(cfg, train_state)
update_actor = make_actor_update_fn(cfg, train_state)
evaluate = make_evaluate_fn(cfg, eval_envs)
if cfg.platform.compile:
mode = "max-autotune-no-cudagraphs"
update_critic = torch.compile(update_critic, mode=mode)
update_actor = torch.compile(update_actor, mode=mode)
postprocess_fn = torch.compile(postprocess_fn, mode=mode)
train_state.compile()
# TODO: Support checkpoint loading
# if cfg.checkpoint_path:
# # Load checkpoint if specified
# torch_checkpoint = torch.load(
# f"{cfg.checkpoint_path}", map_location=device, weights_only=False
# )
# actor.load_state_dict(torch_checkpoint["actor_state_dict"])
# obs_normalizer.load_state_dict(torch_checkpoint["obs_normalizer_state"])
# critic_obs_normalizer.load_state_dict(
# torch_checkpoint["critic_obs_normalizer_state"]
# )
# qnet.load_state_dict(torch_checkpoint["qnet_state_dict"])
# qnet_target.load_state_dict(torch_checkpoint["qnet_target_state_dict"])
# global_step = torch_checkpoint["global_step"]
# else:
global_step = 0
total_env_steps = (
cfg.hyperparameters.total_time_steps
// (cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps)
+ 1
)
pbar = tqdm.tqdm(total=cfg.hyperparameters.total_time_steps, initial=global_step)
start_time = None
desc = ""
eval_interval = total_env_steps // cfg.hyperparameters.num_eval
stochastic_eval = cfg.env.get("stochastic_eval", False)
while global_step < total_env_steps:
if start_time is None and global_step >= cfg.measure_burnin:
start_time = time.time()
measure_burnin = global_step
train_state, transition, infos = collect_fn(train_state)
data = postprocess_fn(train_state, transition)
for _ in range(cfg.hyperparameters.num_epochs):
indices = torch.randperm(
cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps,
device=device,
)
data = data[indices].contiguous()
for j in range(num_batches):
mini_batch = data[j * batch_size : (j + 1) * batch_size]
critic_logs_dict = update_critic(mini_batch)
actor_logs_dict = update_actor(mini_batch)
logs_dict = {
**critic_logs_dict,
**actor_logs_dict,
}
for param, target_param in zip(actor.parameters(), old_actor.parameters()):
target_param.data.copy_(param.data)
if start_time is not None:
# @TODO: shouldn't that be env_steps per second?
speed = (
cfg.hyperparameters.num_envs
* cfg.hyperparameters.num_steps
* (global_step - measure_burnin)
/ (time.time() - start_time)
)
pbar.set_description(f"{speed: 4.4f} sps, " + desc)
with torch.inference_mode():
logs = {
"critic/qf_loss": logs_dict["qf_loss"].mean(),
"critic/qf_max": logs_dict["qf_max"].mean(),
"critic/qf_min": logs_dict["qf_min"].mean(),
"critic/qf_mean": logs_dict["qf_mean"].mean(),
"critic/embedding_loss": logs_dict["embedding_loss"].mean(),
"critic/critic_grad_norm": logs_dict["critic_grad_norm"].mean(),
"actor/actor_loss": logs_dict["actor_loss"].mean(),
"actor/actor_grad_norm": logs_dict["actor_grad_norm"].mean(),
"actor/kl": logs_dict["kl"].mean(),
"actor/entropy": logs_dict["entropy"].mean(),
"actor/temperature": logs_dict["temperature"].mean(),
"actor/lagrangian": logs_dict["lagrangian"].mean(),
"actor/entropy_loss": logs_dict["entropy_loss"].mean(),
"actor/lagrangian_loss": logs_dict["lagrangian_loss"].mean(),
"train/rewards_batch": data["rewards"].mean(),
}
if cfg.env.type == "maniskill":
logs.update(
{
"train/return": torch.stack(
[info["log_info"]["return"] for info in infos]
).mean(),
"train/episode_len": torch.stack(
[info["log_info"]["episode_len"] for info in infos]
)
.float()
.mean(),
"train/success": torch.stack(
[info["log_info"]["success"] for info in infos]
)
.float()
.mean(),
}
)
if eval_interval > 0 and global_step % eval_interval == 0:
print(f"Evaluating at global step {global_step}")
if stochastic_eval:
eval_avg_return, eval_avg_length, stoch_eval_info = evaluate(
train_state, stochastic_eval=stochastic_eval
)
eval_avg_return, eval_avg_length, eval_info = evaluate(
train_state
)
eval_info = {
**eval_info,
**{f"stoch/{k}": v for k, v in stoch_eval_info.items()},
}
else:
eval_avg_return, eval_avg_length, eval_info = evaluate(
train_state
)
if cfg.env.type in [
"humanoid_bench",
"isaaclab",
"mtbench",
]:
# NOTE: Hacky way of evaluating performance, but just works
obs, _ = envs.reset()
logs["eval/avg_return"] = eval_avg_return
logs["eval/avg_length"] = eval_avg_length
for key, value in eval_info.items():
if isinstance(value, torch.Tensor):
logs[f"eval/{key}"] = value.mean().item()
elif isinstance(value, np.ndarray):
logs[f"eval/{key}"] = value.mean()
else:
logs[f"eval/{key}"] = value
print(
f"Eval return: {eval_avg_return:.2f}, length: {eval_avg_length:.2f}, env steps: {global_step * cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps} success rate: {eval_info.get('success', 0.0):.2f}"
)
wandb.log(
{
"speed": speed,
"frame": global_step
* cfg.hyperparameters.num_envs
* cfg.hyperparameters.num_steps,
**logs,
},
step=global_step
* cfg.hyperparameters.num_envs
* cfg.hyperparameters.num_steps,
)
global_step += 1
pbar.update(n=cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps)
if __name__ == "__main__":
main()
+777
View File
@@ -0,0 +1,777 @@
import os
import torch
import torch.nn as nn
from tensordict import TensorDict
class SimpleReplayBuffer(nn.Module):
def __init__(
self,
n_env: int,
buffer_size: int,
n_obs: int,
n_act: int,
n_critic_obs: int,
asymmetric_obs: bool = False,
playground_mode: bool = False,
n_steps: int = 1,
gamma: float = 0.99,
device=None,
):
"""
A simple replay buffer that stores transitions in a circular buffer.
Supports n-step returns and asymmetric observations.
When playground_mode=True, critic_observations are treated as a concatenation of
regular observations and privileged observations, and only the privileged part is stored
to save memory.
TODO (Younggyo): Refactor to split this into SimpleReplayBuffer and NStepReplayBuffer
"""
super().__init__()
self.n_env = n_env
self.buffer_size = buffer_size
self.n_obs = n_obs
self.n_act = n_act
self.n_critic_obs = n_critic_obs
self.asymmetric_obs = asymmetric_obs
self.playground_mode = playground_mode and asymmetric_obs
self.gamma = gamma
self.n_steps = n_steps
self.device = device
self.observations = torch.zeros(
(n_env, buffer_size, n_obs), device=device, dtype=torch.float
)
self.actions = torch.zeros(
(n_env, buffer_size, n_act), device=device, dtype=torch.float
)
self.rewards = torch.zeros(
(n_env, buffer_size), device=device, dtype=torch.float
)
self.dones = torch.zeros((n_env, buffer_size), device=device, dtype=torch.long)
self.truncations = torch.zeros(
(n_env, buffer_size), device=device, dtype=torch.long
)
self.next_observations = torch.zeros(
(n_env, buffer_size, n_obs), device=device, dtype=torch.float
)
if asymmetric_obs:
if self.playground_mode:
# Only store the privileged part of observations (n_critic_obs - n_obs)
self.privileged_obs_size = n_critic_obs - n_obs
self.privileged_observations = torch.zeros(
(n_env, buffer_size, self.privileged_obs_size),
device=device,
dtype=torch.float,
)
self.next_privileged_observations = torch.zeros(
(n_env, buffer_size, self.privileged_obs_size),
device=device,
dtype=torch.float,
)
else:
# Store full critic observations
self.critic_observations = torch.zeros(
(n_env, buffer_size, n_critic_obs), device=device, dtype=torch.float
)
self.next_critic_observations = torch.zeros(
(n_env, buffer_size, n_critic_obs), device=device, dtype=torch.float
)
self.ptr = 0
def extend(
self,
tensor_dict: TensorDict,
):
observations = tensor_dict["observations"]
actions = tensor_dict["actions"]
rewards = tensor_dict["next"]["rewards"]
dones = tensor_dict["next"]["dones"]
truncations = tensor_dict["next"]["truncations"]
next_observations = tensor_dict["next"]["observations"]
ptr = self.ptr % self.buffer_size
self.observations[:, ptr] = observations
self.actions[:, ptr] = actions
self.rewards[:, ptr] = rewards
self.dones[:, ptr] = dones
self.truncations[:, ptr] = truncations
self.next_observations[:, ptr] = next_observations
if self.asymmetric_obs:
critic_observations = tensor_dict["critic_observations"]
next_critic_observations = tensor_dict["next"]["critic_observations"]
if self.playground_mode:
# Extract and store only the privileged part
privileged_observations = critic_observations[:, self.n_obs :]
next_privileged_observations = next_critic_observations[:, self.n_obs :]
self.privileged_observations[:, ptr] = privileged_observations
self.next_privileged_observations[:, ptr] = next_privileged_observations
else:
# Store full critic observations
self.critic_observations[:, ptr] = critic_observations
self.next_critic_observations[:, ptr] = next_critic_observations
self.ptr += 1
def sample(self, batch_size: int):
# we will sample n_env * batch_size transitions
if self.n_steps == 1:
indices = torch.randint(
0,
min(self.buffer_size, self.ptr),
(self.n_env, batch_size),
device=self.device,
)
obs_indices = indices.unsqueeze(-1).expand(-1, -1, self.n_obs)
act_indices = indices.unsqueeze(-1).expand(-1, -1, self.n_act)
observations = torch.gather(self.observations, 1, obs_indices).reshape(
self.n_env * batch_size, self.n_obs
)
next_observations = torch.gather(
self.next_observations, 1, obs_indices
).reshape(self.n_env * batch_size, self.n_obs)
actions = torch.gather(self.actions, 1, act_indices).reshape(
self.n_env * batch_size, self.n_act
)
rewards = torch.gather(self.rewards, 1, indices).reshape(
self.n_env * batch_size
)
dones = torch.gather(self.dones, 1, indices).reshape(
self.n_env * batch_size
)
truncations = torch.gather(self.truncations, 1, indices).reshape(
self.n_env * batch_size
)
effective_n_steps = torch.ones_like(dones)
if self.asymmetric_obs:
if self.playground_mode:
# Gather privileged observations
priv_obs_indices = indices.unsqueeze(-1).expand(
-1, -1, self.privileged_obs_size
)
privileged_observations = torch.gather(
self.privileged_observations, 1, priv_obs_indices
).reshape(self.n_env * batch_size, self.privileged_obs_size)
next_privileged_observations = torch.gather(
self.next_privileged_observations, 1, priv_obs_indices
).reshape(self.n_env * batch_size, self.privileged_obs_size)
# Concatenate with regular observations to form full critic observations
critic_observations = torch.cat(
[observations, privileged_observations], dim=1
)
next_critic_observations = torch.cat(
[next_observations, next_privileged_observations], dim=1
)
else:
# Gather full critic observations
critic_obs_indices = indices.unsqueeze(-1).expand(
-1, -1, self.n_critic_obs
)
critic_observations = torch.gather(
self.critic_observations, 1, critic_obs_indices
).reshape(self.n_env * batch_size, self.n_critic_obs)
next_critic_observations = torch.gather(
self.next_critic_observations, 1, critic_obs_indices
).reshape(self.n_env * batch_size, self.n_critic_obs)
else:
# Sample base indices
if self.ptr >= self.buffer_size:
# When the buffer is full, there is no protection against sampling across different episodes
# We avoid this by temporarily setting self.pos - 1 to truncated = True if not done
# https://github.com/DLR-RM/stable-baselines3/blob/b91050ca94f8bce7a0285c91f85da518d5a26223/stable_baselines3/common/buffers.py#L857-L860
# TODO (Younggyo): Change the reference when this SB3 branch is merged
current_pos = self.ptr % self.buffer_size
curr_truncations = self.truncations[:, current_pos - 1].clone()
self.truncations[:, current_pos - 1] = torch.logical_not(
self.dones[:, current_pos - 1]
)
indices = torch.randint(
0,
self.buffer_size,
(self.n_env, batch_size),
device=self.device,
)
else:
# Buffer not full - ensure n-step sequence doesn't exceed valid data
max_start_idx = max(1, self.ptr - self.n_steps + 1)
indices = torch.randint(
0,
max_start_idx,
(self.n_env, batch_size),
device=self.device,
)
obs_indices = indices.unsqueeze(-1).expand(-1, -1, self.n_obs)
act_indices = indices.unsqueeze(-1).expand(-1, -1, self.n_act)
# Get base transitions
observations = torch.gather(self.observations, 1, obs_indices).reshape(
self.n_env * batch_size, self.n_obs
)
actions = torch.gather(self.actions, 1, act_indices).reshape(
self.n_env * batch_size, self.n_act
)
if self.asymmetric_obs:
if self.playground_mode:
# Gather privileged observations
priv_obs_indices = indices.unsqueeze(-1).expand(
-1, -1, self.privileged_obs_size
)
privileged_observations = torch.gather(
self.privileged_observations, 1, priv_obs_indices
).reshape(self.n_env * batch_size, self.privileged_obs_size)
# Concatenate with regular observations to form full critic observations
critic_observations = torch.cat(
[observations, privileged_observations], dim=1
)
else:
# Gather full critic observations
critic_obs_indices = indices.unsqueeze(-1).expand(
-1, -1, self.n_critic_obs
)
critic_observations = torch.gather(
self.critic_observations, 1, critic_obs_indices
).reshape(self.n_env * batch_size, self.n_critic_obs)
# Create sequential indices for each sample
# This creates a [n_env, batch_size, n_step] tensor of indices
seq_offsets = torch.arange(self.n_steps, device=self.device).view(1, 1, -1)
all_indices = (
indices.unsqueeze(-1) + seq_offsets
) % self.buffer_size # [n_env, batch_size, n_step]
# Gather all rewards and terminal flags
# Using advanced indexing - result shapes: [n_env, batch_size, n_step]
all_rewards = torch.gather(
self.rewards.unsqueeze(-1).expand(-1, -1, self.n_steps), 1, all_indices
)
all_dones = torch.gather(
self.dones.unsqueeze(-1).expand(-1, -1, self.n_steps), 1, all_indices
)
all_truncations = torch.gather(
self.truncations.unsqueeze(-1).expand(-1, -1, self.n_steps),
1,
all_indices,
)
# Create masks for rewards *after* first done
# This creates a cumulative product that zeroes out rewards after the first done
all_dones_shifted = torch.cat(
[torch.zeros_like(all_dones[:, :, :1]), all_dones[:, :, :-1]], dim=2
) # First reward should not be masked
done_masks = torch.cumprod(
1.0 - all_dones_shifted, dim=2
) # [n_env, batch_size, n_step]
effective_n_steps = done_masks.sum(2)
# Create discount factors
discounts = torch.pow(
self.gamma, torch.arange(self.n_steps, device=self.device)
) # [n_steps]
# Apply masks and discounts to rewards
masked_rewards = all_rewards * done_masks # [n_env, batch_size, n_step]
discounted_rewards = masked_rewards * discounts.view(
1, 1, -1
) # [n_env, batch_size, n_step]
# Sum rewards along the n_step dimension
n_step_rewards = discounted_rewards.sum(dim=2) # [n_env, batch_size]
# Find index of first done or truncation or last step for each sequence
first_done = torch.argmax(
(all_dones > 0).float(), dim=2
) # [n_env, batch_size]
first_trunc = torch.argmax(
(all_truncations > 0).float(), dim=2
) # [n_env, batch_size]
# Handle case where there are no dones or truncations
no_dones = all_dones.sum(dim=2) == 0
no_truncs = all_truncations.sum(dim=2) == 0
# When no dones or truncs, use the last index
first_done = torch.where(no_dones, self.n_steps - 1, first_done)
first_trunc = torch.where(no_truncs, self.n_steps - 1, first_trunc)
# Take the minimum (first) of done or truncation
final_indices = torch.minimum(
first_done, first_trunc
) # [n_env, batch_size]
# Create indices to gather the final next observations
final_next_obs_indices = torch.gather(
all_indices, 2, final_indices.unsqueeze(-1)
).squeeze(-1) # [n_env, batch_size]
# Gather final values
final_next_observations = self.next_observations.gather(
1, final_next_obs_indices.unsqueeze(-1).expand(-1, -1, self.n_obs)
)
final_dones = self.dones.gather(1, final_next_obs_indices)
final_truncations = self.truncations.gather(1, final_next_obs_indices)
if self.asymmetric_obs:
if self.playground_mode:
# Gather final privileged observations
final_next_privileged_observations = (
self.next_privileged_observations.gather(
1,
final_next_obs_indices.unsqueeze(-1).expand(
-1, -1, self.privileged_obs_size
),
)
)
# Reshape for output
next_privileged_observations = (
final_next_privileged_observations.reshape(
self.n_env * batch_size, self.privileged_obs_size
)
)
# Concatenate with next observations to form full next critic observations
next_observations_reshaped = final_next_observations.reshape(
self.n_env * batch_size, self.n_obs
)
next_critic_observations = torch.cat(
[next_observations_reshaped, next_privileged_observations],
dim=1,
)
else:
# Gather final next critic observations directly
final_next_critic_observations = (
self.next_critic_observations.gather(
1,
final_next_obs_indices.unsqueeze(-1).expand(
-1, -1, self.n_critic_obs
),
)
)
next_critic_observations = final_next_critic_observations.reshape(
self.n_env * batch_size, self.n_critic_obs
)
# Reshape everything to batch dimension
rewards = n_step_rewards.reshape(self.n_env * batch_size)
dones = final_dones.reshape(self.n_env * batch_size)
truncations = final_truncations.reshape(self.n_env * batch_size)
effective_n_steps = effective_n_steps.reshape(self.n_env * batch_size)
next_observations = final_next_observations.reshape(
self.n_env * batch_size, self.n_obs
)
out = TensorDict(
{
"observations": observations,
"actions": actions,
"next": {
"rewards": rewards,
"dones": dones,
"truncations": truncations,
"observations": next_observations,
"effective_n_steps": effective_n_steps,
},
},
batch_size=self.n_env * batch_size,
)
if self.asymmetric_obs:
out["critic_observations"] = critic_observations
out["next"]["critic_observations"] = next_critic_observations
if self.n_steps > 1 and self.ptr >= self.buffer_size:
# Roll back the truncation flags introduced for safe sampling
self.truncations[:, current_pos - 1] = curr_truncations
return out
class EmpiricalNormalization(nn.Module):
"""Normalize mean and variance of values based on empirical values."""
def __init__(self, shape, device, eps=1e-2, until=None):
"""Initialize EmpiricalNormalization module.
Args:
shape (int or tuple of int): Shape of input values except batch axis.
eps (float): Small value for stability.
until (int or None): If this arg is specified, the link learns input values until the sum of batch sizes
exceeds it.
"""
super().__init__()
self.eps = eps
self.until = until
self.device = device
self.register_buffer("_mean", torch.zeros(shape).unsqueeze(0).to(device))
self.register_buffer("_var", torch.ones(shape).unsqueeze(0).to(device))
self.register_buffer("_std", torch.ones(shape).unsqueeze(0).to(device))
self.register_buffer("count", torch.tensor(0, dtype=torch.long).to(device))
@property
def mean(self):
return self._mean.squeeze(0).clone()
@property
def std(self):
return self._std.squeeze(0).clone()
def forward(self, x: torch.Tensor, center: bool = True) -> torch.Tensor:
if x.shape[-1:] != self._mean.shape[-1:]:
raise ValueError(
f"Expected input of shape (*,{self._mean.shape[-1:]}), got {x.shape}"
)
if self.training:
self.update(x)
if center:
return (x - self._mean) / (self._std + self.eps)
else:
return x / (self._std + self.eps)
@torch.jit.unused
def update(self, x):
x = x.flatten(end_dim=-2)
if self.until is not None and self.count >= self.until:
return
batch_size = x.shape[0]
batch_mean = torch.mean(x, dim=0, keepdim=True)
# Update count
new_count = self.count + batch_size
# Update mean
delta = batch_mean - self._mean
self._mean += (batch_size / new_count) * delta
# Update variance using Chan's parallel algorithm
# https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance#Parallel_algorithm
if self.count > 0: # Ensure we're not dividing by zero
batch_var = torch.mean((x - batch_mean) ** 2, dim=0, keepdim=True)
delta2 = batch_mean - self._mean
m_a = self._var * self.count
m_b = batch_var * batch_size
M2 = m_a + m_b + (delta2**2) * (self.count * batch_size / new_count)
self._var = M2 / new_count
else:
# For first batch, just use batch variance
self._var = torch.mean((x - self._mean) ** 2, dim=0, keepdim=True)
self._std = torch.sqrt(self._var)
self.count = new_count
@torch.jit.unused
def inverse(self, y):
return y * (self._std + self.eps) + self._mean
class RewardNormalizer(nn.Module):
def __init__(
self,
gamma: float,
device: torch.device,
g_max: float = 10.0,
epsilon: float = 1e-8,
):
super().__init__()
self.register_buffer(
"G", torch.zeros(1, device=device)
) # running estimate of the discounted return
self.register_buffer("G_r_max", torch.zeros(1, device=device)) # running-max
self.G_rms = EmpiricalNormalization(shape=1, device=device)
self.gamma = gamma
self.g_max = g_max
self.epsilon = epsilon
def _scale_reward(self, rewards: torch.Tensor) -> torch.Tensor:
var_denominator = self.G_rms.std[0] + self.epsilon
min_required_denominator = self.G_r_max / self.g_max
denominator = torch.maximum(var_denominator, min_required_denominator)
return rewards / denominator
def update_stats(
self,
rewards: torch.Tensor,
dones: torch.Tensor,
):
self.G = self.gamma * (1 - dones) * self.G + rewards
self.G_rms.update(self.G.view(-1, 1))
self.G_r_max = max(self.G_r_max, max(abs(self.G)))
def forward(self, rewards: torch.Tensor) -> torch.Tensor:
return self._scale_reward(rewards)
class PerTaskEmpiricalNormalization(nn.Module):
"""Normalize mean and variance of values based on empirical values for each task."""
def __init__(
self,
num_tasks: int,
shape: tuple,
device: torch.device,
eps: float = 1e-2,
until: int = None,
):
"""
Initialize PerTaskEmpiricalNormalization module.
Args:
num_tasks (int): The total number of tasks.
shape (int or tuple of int): Shape of input values except batch axis.
eps (float): Small value for stability.
until (int or None): If specified, learns until the sum of batch sizes
for a specific task exceeds this value.
"""
super().__init__()
if not isinstance(shape, tuple):
shape = (shape,)
self.num_tasks = num_tasks
self.shape = shape
self.eps = eps
self.until = until
self.device = device
# Buffers now have a leading dimension for tasks
self.register_buffer("_mean", torch.zeros(num_tasks, *shape).to(device))
self.register_buffer("_var", torch.ones(num_tasks, *shape).to(device))
self.register_buffer("_std", torch.ones(num_tasks, *shape).to(device))
self.register_buffer(
"count", torch.zeros(num_tasks, dtype=torch.long).to(device)
)
def forward(
self, x: torch.Tensor, task_ids: torch.Tensor, center: bool = True
) -> torch.Tensor:
"""
Normalize the input tensor `x` using statistics for the given `task_ids`.
Args:
x (torch.Tensor): Input tensor of shape [num_envs, *shape].
task_ids (torch.Tensor): Tensor of task indices, shape [num_envs].
center (bool): If True, center the data by subtracting the mean.
"""
if x.shape[1:] != self.shape:
raise ValueError(f"Expected input shape (*, {self.shape}), got {x.shape}")
if x.shape[0] != task_ids.shape[0]:
raise ValueError("Batch size of x and task_ids must match.")
# Gather the stats for the tasks in the current batch
# Reshape task_ids for broadcasting: [num_envs] -> [num_envs, 1, ...]
view_shape = (task_ids.shape[0],) + (1,) * len(self.shape)
task_ids_expanded = task_ids.view(view_shape).expand_as(x)
mean = self._mean.gather(0, task_ids_expanded)
std = self._std.gather(0, task_ids_expanded)
if self.training:
self.update(x, task_ids)
if center:
return (x - mean) / (std + self.eps)
else:
return x / (std + self.eps)
@torch.jit.unused
def update(self, x: torch.Tensor, task_ids: torch.Tensor):
"""Update running statistics for the tasks present in the batch."""
unique_tasks = torch.unique(task_ids)
for task_id in unique_tasks:
if self.until is not None and self.count[task_id] >= self.until:
continue
# Create a mask to select data for the current task
mask = task_ids == task_id
x_task = x[mask]
batch_size = x_task.shape[0]
if batch_size == 0:
continue
# Update count for this task
old_count = self.count[task_id].clone()
new_count = old_count + batch_size
# Update mean
task_mean = self._mean[task_id]
batch_mean = torch.mean(x_task, dim=0)
delta = batch_mean - task_mean
self._mean[task_id] = task_mean + (batch_size / new_count) * delta
# Update variance using Chan's parallel algorithm
if old_count > 0:
batch_var = torch.var(x_task, dim=0, unbiased=False)
m_a = self._var[task_id] * old_count
m_b = batch_var * batch_size
M2 = m_a + m_b + (delta**2) * (old_count * batch_size / new_count)
self._var[task_id] = M2 / new_count
else:
# For the first batch of this task
self._var[task_id] = torch.var(x_task, dim=0, unbiased=False)
self._std[task_id] = torch.sqrt(self._var[task_id])
self.count[task_id] = new_count
class PerTaskRewardNormalizer(nn.Module):
def __init__(
self,
num_tasks: int,
gamma: float,
device: torch.device,
g_max: float = 10.0,
epsilon: float = 1e-8,
):
"""
Per-task reward normalizer, motivation comes from BRC (https://arxiv.org/abs/2505.23150v1)
"""
super().__init__()
self.num_tasks = num_tasks
self.gamma = gamma
self.g_max = g_max
self.epsilon = epsilon
self.device = device
# Per-task running estimate of the discounted return
self.register_buffer("G", torch.zeros(num_tasks, device=device))
# Per-task running-max of the discounted return
self.register_buffer("G_r_max", torch.zeros(num_tasks, device=device))
# Use the new per-task normalizer for the statistics of G
self.G_rms = PerTaskEmpiricalNormalization(
num_tasks=num_tasks, shape=(1,), device=device
)
def _scale_reward(
self, rewards: torch.Tensor, task_ids: torch.Tensor
) -> torch.Tensor:
"""
Scales rewards using per-task statistics.
Args:
rewards (torch.Tensor): Reward tensor, shape [num_envs].
task_ids (torch.Tensor): Task indices, shape [num_envs].
"""
# Gather stats for the tasks in the batch
std_for_batch = self.G_rms._std.gather(0, task_ids.unsqueeze(-1)).squeeze(-1)
g_r_max_for_batch = self.G_r_max.gather(0, task_ids)
var_denominator = std_for_batch + self.epsilon
min_required_denominator = g_r_max_for_batch / self.g_max
denominator = torch.maximum(var_denominator, min_required_denominator)
# Add a small epsilon to the final denominator to prevent division by zero
# in case g_r_max is also zero.
return rewards / (denominator + self.epsilon)
def update_stats(
self, rewards: torch.Tensor, dones: torch.Tensor, task_ids: torch.Tensor
):
"""
Updates the running discounted return and its statistics for each task.
Args:
rewards (torch.Tensor): Reward tensor, shape [num_envs].
dones (torch.Tensor): Done tensor, shape [num_envs].
task_ids (torch.Tensor): Task indices, shape [num_envs].
"""
if not (rewards.shape == dones.shape == task_ids.shape):
raise ValueError("rewards, dones, and task_ids must have the same shape.")
# === Update G (running discounted return) ===
# Gather the previous G values for the tasks in the batch
prev_G = self.G.gather(0, task_ids)
# Update G for each environment based on its own reward and done signal
new_G = self.gamma * (1 - dones.float()) * prev_G + rewards
# Scatter the updated G values back to the main buffer
self.G.scatter_(0, task_ids, new_G)
# === Update G_rms (statistics of G) ===
# The update function handles the per-task logic internally
self.G_rms.update(new_G.unsqueeze(-1), task_ids)
# === Update G_r_max (running max of |G|) ===
prev_G_r_max = self.G_r_max.gather(0, task_ids)
# Update the max for each environment
updated_G_r_max = torch.maximum(prev_G_r_max, torch.abs(new_G))
# Scatter the new maxes back to the main buffer
self.G_r_max.scatter_(0, task_ids, updated_G_r_max)
def forward(self, rewards: torch.Tensor, task_ids: torch.Tensor) -> torch.Tensor:
"""
Normalizes rewards. During training, it also updates the running statistics.
Args:
rewards (torch.Tensor): Reward tensor, shape [num_envs].
task_ids (torch.Tensor): Task indices, shape [num_envs].
"""
return self._scale_reward(rewards, task_ids)
def cpu_state(sd):
# detach & move to host without locking the compute stream
return {k: v.detach().to("cpu", non_blocking=True) for k, v in sd.items()}
def save_params(
global_step,
actor,
qnet,
qnet_target,
obs_normalizer,
critic_obs_normalizer,
args,
save_path,
):
"""Save model parameters and training configuration to disk."""
os.makedirs(os.path.dirname(save_path), exist_ok=True)
save_dict = {
"actor_state_dict": cpu_state(actor.state_dict()),
"qnet_state_dict": cpu_state(qnet.state_dict()),
"qnet_target_state_dict": cpu_state(qnet_target.state_dict()),
"obs_normalizer_state": (
cpu_state(obs_normalizer.state_dict())
if hasattr(obs_normalizer, "state_dict")
else None
),
"critic_obs_normalizer_state": (
cpu_state(critic_obs_normalizer.state_dict())
if hasattr(critic_obs_normalizer, "state_dict")
else None
),
"args": vars(args), # Save all arguments
"global_step": global_step,
}
torch.save(save_dict, save_path, _use_new_zipfile_serialization=True)
print(f"Saved parameters and configuration to {save_path}")
def hl_gauss(inp, vmin, vmax, num_atoms):
x = torch.clip(inp, vmin, max=vmax)
bin_width = (vmax - vmin) / (num_atoms - 1)
sigma_to_final_sigma_ratio = 0.75
support = torch.linspace(
vmin - bin_width / 2,
vmax + bin_width / 2,
num_atoms + 1,
device=inp.device,
)
sigma = bin_width * sigma_to_final_sigma_ratio
cdf_evals = torch.erf(
(support.unsqueeze(0) - x).squeeze()
/ (torch.sqrt(torch.tensor(2.0)) * sigma + 1e-6)
)
z = cdf_evals[..., -1] - cdf_evals[..., 0]
target_probs = cdf_evals[..., 1:] - cdf_evals[..., :-1]
target_probs = (target_probs / (z.unsqueeze(-1) + 1e-6)).reshape(
*inp.shape[:-1], num_atoms
)
return target_probs