Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions ss2r/algorithms/mbpo/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
"wandb_id",
"hard_resets",
"nonepisodic",
"action_delay",
]
}
policy_hidden_layer_sizes = agent_cfg.pop("policy_hidden_layer_sizes")
Expand Down
1 change: 1 addition & 0 deletions ss2r/algorithms/ppo/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
"wandb_id",
"hard_resets",
"nonepisodic",
"action_delay",
]
}
policy_hidden_layer_sizes = agent_cfg.pop("policy_hidden_layer_sizes")
Expand Down
1 change: 1 addition & 0 deletions ss2r/algorithms/sac/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
"wandb_id",
"hard_resets",
"nonepisodic",
"action_delay",
]
}
policy_hidden_layer_sizes = agent_cfg.pop("policy_hidden_layer_sizes")
Expand Down
1 change: 1 addition & 0 deletions ss2r/algorithms/sbsrl/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
"wandb_id",
"hard_resets",
"nonepisodic",
"action_delay",
]
}
policy_hidden_layer_sizes = agent_cfg.pop("policy_hidden_layer_sizes")
Expand Down
23 changes: 23 additions & 0 deletions ss2r/benchmark_suites/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@
from ss2r.benchmark_suites.safety_gym import go_to_goal
from ss2r.benchmark_suites.utils import get_domain_name, get_task_config
from ss2r.benchmark_suites.wrappers import (
ActionDelayWrapper,
GoToGoalObservationWrapper,
Saute,
SPiDR,
Expand Down Expand Up @@ -385,19 +386,39 @@ def make_spidr_cartpole_vision(cfg, train_wrap_env_fn, eval_wrap_env_fn):
return train_env, train_env


def _get_action_delay_max(task_cfg, training_cfg):
def _extract(cfg_section):
if cfg_section is None or not hasattr(cfg_section, "get"):
return None
if not cfg_section.get("enable", False):
return None
max_delay = int(cfg_section.get("max_delay", 0))
if max_delay <= 0:
return None
return max_delay

max_delay = _extract(task_cfg.get("action_delay"))
if max_delay is not None:
return max_delay
return _extract(training_cfg.get("action_delay"))


def make_mujoco_playground_envs(cfg, train_wrap_env_fn, eval_wrap_env_fn):
from ml_collections import config_dict
from mujoco_playground import registry

from ss2r.benchmark_suites.mujoco_playground import wrap_for_brax_training

task_cfg = get_task_config(cfg)
action_delay_max = _get_action_delay_max(task_cfg, cfg.training)
task_params = config_dict.ConfigDict(task_cfg.task_params)
vision = "use_vision" in cfg.agent and cfg.agent.use_vision
if vision:
_preinitialize_vision_env(task_cfg.task_name, task_params, registry)
train_env = registry.load(task_cfg.task_name, config=task_params)
train_env = train_wrap_env_fn(train_env)
if action_delay_max is not None:
train_env = ActionDelayWrapper(train_env, action_delay_max)
train_key, eval_key = jax.random.split(jax.random.PRNGKey(cfg.training.seed))
if vision and cfg.training.train_domain_randomization:
train_randomization_fn = functools.partial(
Expand Down Expand Up @@ -429,6 +450,8 @@ def make_mujoco_playground_envs(cfg, train_wrap_env_fn, eval_wrap_env_fn):
return train_env, train_env
eval_env = registry.load(task_cfg.task_name, config=task_params)
eval_env = eval_wrap_env_fn(eval_env)
if action_delay_max is not None:
eval_env = ActionDelayWrapper(eval_env, action_delay_max)
eval_randomization_fn = (
prepare_randomization_fn(
eval_key,
Expand Down
32 changes: 32 additions & 0 deletions ss2r/benchmark_suites/wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,38 @@ def _env_fn(self, model):
return env


class ActionDelayWrapper(Wrapper):
def __init__(self, env: Env, max_delay: int):
super().__init__(env)
if max_delay < 0:
raise ValueError("max_delay must be >= 0")
self._max_delay = int(max_delay)

def reset(self, rng: jax.Array) -> State | MjxState:
rng, delay_rng = jax.random.split(rng)
state = self.env.reset(rng)
if self._max_delay == 0:
return state
action_buffer = jp.zeros((self._max_delay + 1, self.action_size))
state.info["action_delay_rng"] = delay_rng
state.info["action_delay_buffer"] = action_buffer
return state

def step(self, state: State | MjxState, action: jax.Array) -> State | MjxState:
if self._max_delay == 0:
return self.env.step(state, action)
rng = state.info["action_delay_rng"]
action_buffer = state.info["action_delay_buffer"]
rng, key = jax.random.split(rng)
action_buffer = jp.roll(action_buffer, shift=-1, axis=0)
action_buffer = action_buffer.at[-1].set(action)
delay = jax.random.randint(key, (), minval=0, maxval=self._max_delay + 1)
delayed_action = action_buffer[self._max_delay - delay]
state.info["action_delay_rng"] = rng
state.info["action_delay_buffer"] = action_buffer
return self.env.step(state, delayed_action)


class CostEpisodeWrapper(brax_training.EpisodeWrapper):
"""Maintains episode step count and sets done at episode end."""

Expand Down
26 changes: 26 additions & 0 deletions ss2r/configs/experiment/go1_sim_to_real_transfer_stability.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
# @package _global_
defaults:
- go1_online
- override /agent/data_collection: episodic
- override /agent/replay_buffer: pytree
- override /environment: go1_joystick
- _self_



training:
train_domain_randomization: false
eval_domain_randomization: true
num_eval_envs: 128
render: false
num_evals: 25
num_timesteps: 25000
wandb_id: u5n93as3
action_delay:
max_delay: 0
enable: false


environment:
task_params:
ctrl_dt: 0.02
5 changes: 4 additions & 1 deletion ss2r/configs/train_brax.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -52,4 +52,7 @@ training:
store_checkpoint: true
wandb_id: null
hard_resets: false
nonepisodic: false
nonepisodic: false
action_delay:
enable: false
max_delay: 0