diff --git a/ss2r/algorithms/mbpo/__init__.py b/ss2r/algorithms/mbpo/__init__.py index fedd30e22..6390f0371 100644 --- a/ss2r/algorithms/mbpo/__init__.py +++ b/ss2r/algorithms/mbpo/__init__.py @@ -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") diff --git a/ss2r/algorithms/ppo/__init__.py b/ss2r/algorithms/ppo/__init__.py index 341a8e095..85ed3df25 100644 --- a/ss2r/algorithms/ppo/__init__.py +++ b/ss2r/algorithms/ppo/__init__.py @@ -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") diff --git a/ss2r/algorithms/sac/__init__.py b/ss2r/algorithms/sac/__init__.py index e536531dc..df5b518d7 100644 --- a/ss2r/algorithms/sac/__init__.py +++ b/ss2r/algorithms/sac/__init__.py @@ -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") diff --git a/ss2r/algorithms/sbsrl/__init__.py b/ss2r/algorithms/sbsrl/__init__.py index d0bf42890..e5696a48e 100644 --- a/ss2r/algorithms/sbsrl/__init__.py +++ b/ss2r/algorithms/sbsrl/__init__.py @@ -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") diff --git a/ss2r/benchmark_suites/__init__.py b/ss2r/benchmark_suites/__init__.py index 3f7cb8631..e7f03092e 100644 --- a/ss2r/benchmark_suites/__init__.py +++ b/ss2r/benchmark_suites/__init__.py @@ -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, @@ -385,6 +386,23 @@ 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 @@ -392,12 +410,15 @@ def make_mujoco_playground_envs(cfg, train_wrap_env_fn, eval_wrap_env_fn): 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( @@ -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, diff --git a/ss2r/benchmark_suites/wrappers.py b/ss2r/benchmark_suites/wrappers.py index 7dfa95ca9..5f5a0229a 100644 --- a/ss2r/benchmark_suites/wrappers.py +++ b/ss2r/benchmark_suites/wrappers.py @@ -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.""" diff --git a/ss2r/configs/experiment/go1_sim_to_real_transfer_stability.yaml b/ss2r/configs/experiment/go1_sim_to_real_transfer_stability.yaml new file mode 100644 index 000000000..e72a22b50 --- /dev/null +++ b/ss2r/configs/experiment/go1_sim_to_real_transfer_stability.yaml @@ -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 \ No newline at end of file diff --git a/ss2r/configs/train_brax.yaml b/ss2r/configs/train_brax.yaml index f1cfaf2e7..2b37ce746 100644 --- a/ss2r/configs/train_brax.yaml +++ b/ss2r/configs/train_brax.yaml @@ -52,4 +52,7 @@ training: store_checkpoint: true wandb_id: null hard_resets: false - nonepisodic: false \ No newline at end of file + nonepisodic: false + action_delay: + enable: false + max_delay: 0