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