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 @@ -24,6 +24,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
"value_privileged",
"policy_privileged",
"wandb_id",
"hard_resets",
]
}
policy_hidden_layer_sizes = agent_cfg.pop("policy_hidden_layer_sizes")
Expand Down
3 changes: 2 additions & 1 deletion ss2r/algorithms/ppo/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,8 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
"eval_domain_randomization",
"render",
"store_checkpoint",
"wandb_id", # Add this to the exclusion list
"wandb_id",
"hard_resets",
]
}
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 @@ -28,6 +28,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
"value_privileged",
"policy_privileged",
"wandb_id",
"hard_resets",
]
}
policy_hidden_layer_sizes = agent_cfg.pop("policy_hidden_layer_sizes")
Expand Down
3 changes: 3 additions & 0 deletions ss2r/benchmark_suites/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,6 +186,7 @@ def make_brax_envs(cfg, train_wrap_env_fn, eval_wrap_env_fn):
action_repeat=cfg.training.action_repeat,
randomization_fn=train_randomization_fn,
augment_state=False,
hard_resets=cfg.training.hard_resets,
)
eval_randomization_fn = prepare_randomization_fn(
eval_key, cfg.training.num_eval_envs, task_cfg.eval_params, task_cfg.task_name
Expand Down Expand Up @@ -228,6 +229,7 @@ def make_mujoco_playground_envs(cfg, train_wrap_env_fn, eval_wrap_env_fn):
episode_length=cfg.training.episode_length,
action_repeat=cfg.training.action_repeat,
augment_state=False,
hard_resets=cfg.training.hard_resets,
)
eval_randomization_fn = (
prepare_randomization_fn(
Expand Down Expand Up @@ -271,6 +273,7 @@ def make_safety_gym_envs(cfg, train_wrap_env_fn, eval_wrap_env_fn):
episode_length=cfg.training.episode_length,
action_repeat=cfg.training.action_repeat,
augment_state=False,
hard_resets=cfg.training.hard_resets,
)
eval_randomization_fn = (
prepare_randomization_fn(
Expand Down
6 changes: 5 additions & 1 deletion ss2r/benchmark_suites/mujoco_playground/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,7 @@ def wrap_for_brax_training(
randomization_fn: Optional[
Callable[[mjx.Model], Tuple[mjx.Model, mjx.Model]]
] = None,
hard_resets: bool = False,
*,
augment_state: bool = False,
) -> mujoco_playground_wrapper.Wrapper:
Expand Down Expand Up @@ -87,5 +88,8 @@ def wrap_for_brax_training(
env, randomization_fn, augment_state=augment_state
)
env = wrappers.CostEpisodeWrapper(env, episode_length, action_repeat)
env = mujoco_playground_wrapper.BraxAutoResetWrapper(env)
if hard_resets:
env = wrappers.HardAutoResetWrapper(env)
else:
env = mujoco_playground_wrapper.BraxAutoResetWrapper(env)
return env
54 changes: 52 additions & 2 deletions ss2r/benchmark_suites/wrappers.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from brax.envs.base import Env, State, Wrapper
from brax.envs.wrappers import training as brax_training
from jax import numpy as jp
from mujoco_playground import State as MjxState


class ActionObservationDelayWrapper(Wrapper):
Expand Down Expand Up @@ -297,6 +298,8 @@ def wrap(
randomization_fn: Optional[
Callable[[System], Tuple[System, System, jax.Array]]
] = None,
hard_resets: bool = False,
*,
augment_state: bool = True,
) -> Wrapper:
"""Common wrapper pattern for all training agents.
Expand All @@ -313,14 +316,17 @@ def wrap(
environment did not already have batch dimensions, it is additional Vmap
wrapped.
"""
env = CostEpisodeWrapper(env, episode_length, action_repeat)
if randomization_fn is None:
env = brax_training.VmapWrapper(env)
else:
env = DomainRandomizationVmapWrapper(
env, randomization_fn, augment_state=augment_state
)
env = brax_training.AutoResetWrapper(env)
env = CostEpisodeWrapper(env, episode_length, action_repeat)
if hard_resets:
env = HardAutoResetWrapper(env)
else:
env = brax_training.AutoResetWrapper(env)
return env


Expand Down Expand Up @@ -400,3 +406,47 @@ def _env_fn(self, model):
env = self.env
env.unwrapped._mjx_model = model
return env


class HardAutoResetWrapper(Wrapper):
"""Automatically reset Brax envs that are done.

Resample only when >=1 environment is actually done. Still resamples for all
"""

def reset(self, rng: jax.Array) -> State | MjxState:
rng, sample_rng = jax.vmap(jax.random.split, out_axes=1)(rng)
state = self.env.reset(sample_rng)
state.info["reset_rng"] = rng
return state

def step(self, state: State | MjxState, action: jax.Array) -> State | MjxState:
if "steps" in state.info:
steps = state.info["steps"]
steps = jp.where(state.done, jp.zeros_like(steps), steps)
state.info.update(steps=steps)
state = state.replace(done=jp.zeros_like(state.done))
state = self.env.step(state, action)
maybe_reset = jax.lax.cond(
state.done.any(), self.reset, lambda rng: state, state.info["reset_rng"]
)

def where_done(x, y):
done = state.done
if done.shape:
done = jp.reshape(done, [x.shape[0]] + [1] * (len(x.shape) - 1)) # type: ignore
return jp.where(done, x, y)

if hasattr(state, "pipeline_state"):
state_data = state.pipeline_state
maybe_reset_data = maybe_reset.pipeline_state
data_name = "pipeline_state"
elif hasattr(state, "data"):
state_data = state.data
maybe_reset_data = maybe_reset.data
data_name = "data"
else:
raise NotImplementedError
new_data = jax.tree.map(where_done, maybe_reset_data, state_data)
obs = jax.tree.map(maybe_reset.obs, state.obs)
return state.replace(**{data_name: new_data, "obs": obs})
3 changes: 2 additions & 1 deletion ss2r/configs/train_brax.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -49,4 +49,5 @@ training:
train_domain_randomization: true
eval_domain_randomization: true
store_checkpoint: true
wandb_id: null
wandb_id: null
hard_resets: false