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
5 changes: 5 additions & 0 deletions ss2r/benchmark_suites/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from ss2r.benchmark_suites.mujoco_playground.cartpole.spidr_cartpole import (
VisionSPiDRCartpole,
)
from ss2r.benchmark_suites.mujoco_playground.g1_joystick import g1_joystick
from ss2r.benchmark_suites.mujoco_playground.go1_joystick import go1_joystick
from ss2r.benchmark_suites.mujoco_playground.go2_joystick import (
getup,
Expand Down Expand Up @@ -499,6 +500,7 @@ def make_safety_gym_envs(cfg, train_wrap_env_fn, eval_wrap_env_fn):
"rccar": rccar.domain_randomization,
"humanoid": humanoid.domain_randomization,
"humanoid_safe": humanoid.domain_randomization,
"G1JoystickFlatTerrain": g1_joystick.domain_randomization,
"Go1JoystickFlatTerrain": go1_joystick.domain_randomization,
"Go1JoystickRoughTerrain": go1_joystick.domain_randomization,
"SafeJointGo1JoystickFlatTerrain": go1_joystick.domain_randomization,
Expand Down Expand Up @@ -549,6 +551,9 @@ def make_safety_gym_envs(cfg, train_wrap_env_fn, eval_wrap_env_fn):
"ant": functools.partial(brax.render, camera="track"),
"ant_safe": functools.partial(brax.render, camera="track"),
"rccar": rccar.render,
"G1JoystickFlatTerrain": functools.partial(
mujoco_playground.render, camera="track"
),
"Go1JoystickFlatTerrain": functools.partial(
mujoco_playground.render, camera="track"
),
Expand Down
Empty file.
127 changes: 127 additions & 0 deletions ss2r/benchmark_suites/mujoco_playground/g1_joystick/g1_joystick.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,127 @@
"""Utilities for randomization."""

import jax
import jax.numpy as jnp
from mujoco import mjx

FLOOR_GEOM_ID = 0
TORSO_BODY_ID = 16
NUM_ACTUATED_DOFS = 29


def domain_randomization(model: mjx.Model, rng: jax.Array, cfg):
@jax.vmap
def rand_dynamics(rng):
# Floor / foot friction: =U(0.4, 1.0).
rng, key = jax.random.split(rng)
friction = jax.random.uniform(
key, minval=cfg.floor_friction[0], maxval=cfg.floor_friction[1]
)
pair_friction = model.pair_friction.at[
FLOOR_GEOM_ID : FLOOR_GEOM_ID + 2,
FLOOR_GEOM_ID : FLOOR_GEOM_ID + 2,
].set(friction)

# Scale static friction: *U(cfg.scale_friction).
rng, key = jax.random.split(rng)
friction_scale = jax.random.uniform(
key,
shape=(NUM_ACTUATED_DOFS,),
minval=cfg.scale_friction[0],
maxval=cfg.scale_friction[1],
)
frictionloss = model.dof_frictionloss[6:] * friction_scale
dof_frictionloss = model.dof_frictionloss.at[6:].set(frictionloss)

# Scale armature: *U(1.0, 1.05).
rng, key = jax.random.split(rng)
armature_scale = jax.random.uniform(
key,
shape=(NUM_ACTUATED_DOFS,),
minval=cfg.scale_armature[0],
maxval=cfg.scale_armature[1],
)
armature = model.dof_armature[6:] * armature_scale
dof_armature = model.dof_armature.at[6:].set(armature)

# Scale all link masses: *U(0.9, 1.1).
rng, key = jax.random.split(rng)
dmass = jax.random.uniform(
key,
shape=(model.nbody,),
minval=cfg.scale_link_mass[0],
maxval=cfg.scale_link_mass[1],
)
body_mass = model.body_mass.at[:].set(model.body_mass * dmass)

# Add mass to torso: +U(-1.0, 1.0).
rng, key = jax.random.split(rng)
dmass_torso = jax.random.uniform(
key, minval=cfg.add_torso_mass[0], maxval=cfg.add_torso_mass[1]
)
body_mass = body_mass.at[TORSO_BODY_ID].set(
body_mass[TORSO_BODY_ID] + dmass_torso
)

# Jitter qpos0: +U(-0.05, 0.05).
rng, key = jax.random.split(rng)
qpos0_jitter = jax.random.uniform(
key,
shape=(NUM_ACTUATED_DOFS,),
minval=cfg.jitter_qpos0[0],
maxval=cfg.jitter_qpos0[1],
)
qpos0 = model.qpos0
qpos0 = qpos0.at[7:].set(qpos0[7:] + qpos0_jitter)

samples = jnp.hstack(
[
jnp.array([friction]),
friction_scale,
armature_scale,
dmass,
jnp.array([dmass_torso]),
qpos0_jitter,
]
)

return (
pair_friction,
dof_frictionloss,
dof_armature,
body_mass,
qpos0,
samples,
)

(
pair_friction,
dof_frictionloss,
dof_armature,
body_mass,
qpos0,
samples,
) = rand_dynamics(rng)

in_axes = jax.tree_util.tree_map(lambda x: None, model)
in_axes = in_axes.tree_replace(
{
"pair_friction": 0,
"dof_frictionloss": 0,
"dof_armature": 0,
"body_mass": 0,
"qpos0": 0,
}
)

model = model.tree_replace(
{
"pair_friction": pair_friction,
"dof_frictionloss": dof_frictionloss,
"dof_armature": dof_armature,
"body_mass": body_mass,
"qpos0": qpos0,
}
)

return model, in_axes, samples
104 changes: 104 additions & 0 deletions ss2r/configs/environment/g1_joystick.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
defaults:
- mujoco_playground_base
task_name: G1JoystickFlatTerrain
task_params:
torque_limit: 12.
Kd: 0.5
Kp: 35.0
ctrl_dt: 0.02
sim_dt: 0.002
episode_length: 1000
action_repeat: 1
action_scale: 0.5
history_len: 1
restricted_joint_range: false
soft_joint_pos_limit_factor: 0.95
noise_config:
level: 1.0
scales:
joint_pos: 0.03
joint_vel: 1.5
gravity: 0.05
linvel: 0.1
gyro: 0.2
reward_config:
scales:
tracking_lin_vel: 1.0
tracking_ang_vel: 0.75
lin_vel_z: 0.0
ang_vel_xy: -0.15
orientation: -2.0
base_height: 0.0
torques: 0.0
action_rate: 0.0
energy: 0.0
dof_acc: 0.0
feet_clearance: 0.0
feet_air_time: 2.0
feet_slip: -0.25
feet_height: 0.0
feet_phase: 1.0
alive: 0.0
stand_still: -1.0
termination: -100.0
collision: -0.1
contact_force: -0.01
joint_deviation_knee: -0.1
joint_deviation_hip: -0.25
dof_pos_limits: -1.0
pose: -0.1
tracking_sigma: 0.25
max_foot_height: 0.15
base_height_target: 0.5
max_contact_force: 500.0
push_config:
enable: true
interval_range:
- 5.0
- 10.0
magnitude_range:
- 0.1
- 2.0
command_config:
a:
- 1.0
- 0.8
- 1.0
b:
- 0.9
- 0.25
- 0.5
lin_vel_x:
- -1.0
- 1.0
lin_vel_y:
- -0.5
- 0.5
ang_vel_yaw:
- -1.0
- 1.0
impl: jax
nconmax: 65536
njmax: 90

train_params:
floor_friction: [0.4, 1.0]
scale_friction: [0.9, 1.1]
scale_armature: [1.0, 1.05]
jitter_mass: [-0.05, 0.05]
scale_link_mass: [0.9, 1.1]
add_torso_mass: [-1.0, 1.0]
jitter_qpos0: [-0.05, 0.05]
Kd: [0.0, 0.0]
Kp: [0.0, 0.0]

eval_params:
floor_friction: [0.4, 1.0]
scale_friction: [0.9, 1.1]
scale_armature: [1.0, 1.05]
jitter_mass: [-0.05, 0.05]
scale_link_mass: [0.9, 1.1]
add_torso_mass: [-1.0, 1.0]
jitter_qpos0: [-0.05, 0.05]
Kd: [0.0, 0.0]
Kp: [0.0, 0.0]
10 changes: 10 additions & 0 deletions ss2r/configs/experiment/g1_sim_to_real_sac.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
# @package _global_
defaults:
- g1_joystick
- override /agent/penalizer: lagrangian
- _self_

training:
train_domain_randomization: true
eval_domain_randomization: true
safe: false