unitree-g1-mujoco / sim /base_sim.py
k-valentin's picture
feat: attach-on-close kinematic grasp for the pick_cylinder world (WORLD_GRASP)
43d0244
Raw History Blame Contribute Delete
42.8 kB
import argparse
import logging
import pathlib
from pathlib import Path
import threading
from threading import Thread
from typing import Dict
import mujoco
import mujoco.viewer
import numpy as np
try:
import rclpy
HAS_RCLPY = True
except ImportError:
HAS_RCLPY = False
print("ROS 2 integration unavailable; camera images use the ZMQ publisher.")
from unitree_sdk2py.core.channel import ChannelFactoryInitialize
import yaml
import os
from .image_publish_utils import ImagePublishProcess
from .keyboard_forward import forward_key
from .metric_utils import check_contact
from .sim_utils import get_subtree_body_names
from .unitree_sdk2py_bridge import ElasticBand, UnitreeSdk2Bridge
from .model_config import select_end_effector
# AmazingHand servo targets (deg) of the open hand, as lerobot's AMAZING_HAND open_q: (servo 1, servo 2) per finger.
AMAZING_HAND_OPEN_DEG = (-35.0, 35.0)
WORLD_EDGE_MARGIN = 0.05 # m from a world object's edge to the table's edge, whatever WORLD_RANDOMIZE
AMAZING_HAND_CLOSED_DEG = (60.0, -60.0)
# Kinematic grasp: closure thresholds and the zone (min, max corners, right hand) in the {side}_tcp frame between
# the palm face and the closed fingertips; the left hand mirrors y.
GRASP_ATTACH_CLOSURE = 0.6
GRASP_RELEASE_CLOSURE = 0.4
GRASP_ZONE = (np.array([0.03, 0.0, -0.07]), np.array([0.11, 0.07, 0.05]))
logger = logging.getLogger(__name__)
GR00T_WBC_ROOT = Path(__file__).resolve().parent.parent # Points to mujoco_sim_g1/
def amazing_hand_closure(servo_q) -> float:
"""Closure 0..1 of an AmazingHand's fingers 1-3 from their six servo angles (rad), as lerobot's q_to_closure."""
q = np.degrees(np.asarray(servo_q, float))
opened = np.tile(AMAZING_HAND_OPEN_DEG, 3)
closed = np.tile(AMAZING_HAND_CLOSED_DEG, 3)
return float(np.clip(np.mean((q - opened) / (closed - opened)), 0.0, 1.0))
class DefaultEnv:
"""Base environment class that handles simulation environment setup and step"""
def __init__(
self,
config: Dict[str, any],
env_name: str = "default",
camera_configs: Dict[str, any] = None,
onscreen: bool = False,
offscreen: bool = False,
):
# Avoid mutable default argument gotcha
if camera_configs is None:
camera_configs = {}
# global_view is only set up for this specifc scene for now.
if config["ROBOT_SCENE"] == "gr00t_wbc/control/robot_model/model_data/g1/scene_29dof.xml":
camera_configs["global_view"] = {
"height": 400,
"width": 400,
}
self.config = config
self.env_name = env_name
self.num_body_dof = self.config["NUM_JOINTS"]
self.num_hand_dof = self.config["NUM_HAND_JOINTS"]
self.sim_dt = self.config["SIMULATE_DT"]
self.obs = None
self.torque_limit = np.array(self.config["motor_effort_limit_list"])
self.camera_configs = camera_configs
# Debug: print camera config
if len(camera_configs) > 0:
print(f"✓ DefaultEnv initialized with {len(camera_configs)} camera(s): {list(camera_configs.keys())}")
# Unitree bridge will be initialized by the simulator
self.unitree_bridge = None
# Store display mode
self.onscreen = onscreen
# Initialize scene (defined in subclasses)
self.init_scene()
# Setup offscreen rendering if needed (lazy init - renderers created on first use)
self.offscreen = offscreen
self.renderers = {} # Will be lazily initialized
self._renderers_initialized = False
self.image_dt = self.config.get("IMAGE_DT", 0.033333)
# Image publishing subprocess (initialized separately)
self.image_publish_process = None
def init_scene(self):
"""Initialize the default robot scene"""
assets_root = Path(__file__).parent.parent
self.mj_model = mujoco.MjModel.from_xml_path(
str(assets_root / self.config["ROBOT_SCENE"])
)
self.mj_data = mujoco.MjData(self.mj_model)
# Set valid floating base quaternion (MjData initializes qpos to zeros)
self.mj_data.qpos[3:7] = [1.0, 0.0, 0.0, 0.0]
self.mj_model.opt.timestep = self.sim_dt
self.torso_index = mujoco.mj_name2id(self.mj_model, mujoco.mjtObj.mjOBJ_BODY, "torso_link")
self.root_body = "pelvis"
# Enable the elastic band
if self.config["ENABLE_ELASTIC_BAND"]:
self.elastic_band = ElasticBand()
if "g1" in self.config["ROBOT_TYPE"]:
if self.config["enable_waist"]:
self.band_attached_link = self.mj_model.body("pelvis").id
else:
self.band_attached_link = self.mj_model.body("torso_link").id
elif "h1" in self.config["ROBOT_TYPE"]:
self.band_attached_link = self.mj_model.body("torso_link").id
else:
self.band_attached_link = self.mj_model.body("base_link").id
if self.onscreen:
self.viewer = mujoco.viewer.launch_passive(
self.mj_model,
self.mj_data,
key_callback=self._viewer_key_callback,
show_left_ui=False,
show_right_ui=False,
)
else:
mujoco.mj_forward(self.mj_model, self.mj_data)
self.viewer = None
else:
if self.onscreen:
self.viewer = mujoco.viewer.launch_passive(
self.mj_model, self.mj_data, show_left_ui=False, show_right_ui=False
)
else:
mujoco.mj_forward(self.mj_model, self.mj_data)
self.viewer = None
if self.viewer:
# viewer camera
self.viewer.cam.azimuth = 120 # Horizontal rotation in degrees
self.viewer.cam.elevation = -30 # Vertical tilt in degrees
self.viewer.cam.distance = 2.0 # Distance from camera to target
self.viewer.cam.lookat = np.array([0, 0, 0.5]) # Point the camera is looking at
# Body DDS indices exclude the end effectors. Map each scalar joint
# explicitly: actuator order need not match the model's joint order.
# AmazingHand fingers are not DDS-managed like dex1/dex3 (a later change wires them
# to their own control path), so they are excluded from this DDS hand scan entirely.
is_dds_hand = self.config.get("END_EFFECTOR") != "amazing_hand"
actuated_joint_ids = set(self.mj_model.actuator_trnid[
self.mj_model.actuator_trntype == mujoco.mjtTrn.mjTRN_JOINT, 0
].tolist())
self.body_joint_index = []
self.left_hand_index = []
self.right_hand_index = []
for i in range(self.mj_model.njnt):
name = self.mj_model.joint(i).name
if any(
[
part_name in name
for part_name in ["hip", "knee", "ankle", "waist", "shoulder", "elbow", "wrist"]
]
):
self.body_joint_index.append(i)
elif not is_dds_hand or i not in actuated_joint_ids:
continue
elif "left_hand" in name or name.startswith("left_dex1_finger_joint_"):
self.left_hand_index.append(i)
elif "right_hand" in name or name.startswith("right_dex1_finger_joint_"):
self.right_hand_index.append(i)
assert len(self.body_joint_index) == self.config["NUM_JOINTS"], \
f"Expected {self.config['NUM_JOINTS']} body joints, got {len(self.body_joint_index)}"
expected_hands = self.config.get("NUM_HAND_JOINTS", 0)
if len(self.left_hand_index) != expected_hands or len(self.right_hand_index) != expected_hands:
raise ValueError(f"Expected {expected_hands} joints per end effector, got left={len(self.left_hand_index)}, right={len(self.right_hand_index)}")
for prefix, attribute in (("body", "body_joint_index"), ("left_hand", "left_hand_index"), ("right_hand", "right_hand_index")):
joint_ids = np.asarray(getattr(self, attribute), dtype=int)
setattr(self, attribute, joint_ids)
setattr(self, prefix + "_qpos_index", self.mj_model.jnt_qposadr[joint_ids])
setattr(self, prefix + "_dof_index", self.mj_model.jnt_dofadr[joint_ids])
actuator_ids = []
for joint_id in joint_ids:
matches = np.flatnonzero(
(self.mj_model.actuator_trntype == mujoco.mjtTrn.mjTRN_JOINT)
& (self.mj_model.actuator_trnid[:, 0] == joint_id)
)
if len(matches) != 1:
raise ValueError(f"Expected one actuator for {self.mj_model.joint(joint_id).name}, got {len(matches)}")
actuator_ids.append(matches[0])
setattr(self, prefix + "_actuator_index", np.asarray(actuator_ids, dtype=int))
self.dds_actuator_index = np.concatenate(
(self.body_actuator_index, self.left_hand_actuator_index, self.right_hand_actuator_index)
).astype(int)
self.torques = np.zeros(self.mj_model.nu)
if self.config.get("FREE_BASE", False):
self.torque_limit = np.concatenate((np.zeros(6), self.torque_limit))
if self.torque_limit.shape[0] < self.torques.shape[0]:
# motor_effort_limit_list only covers the DDS-managed actuators (body + dex1/dex3
# hands). Non-DDS actuators (e.g. the D455 pan/tilt head, AmazingHand fingers) are
# never written from self.torques (see dds_actuator_index below), so their clip
# bound is a don't-care; pad with +inf rather than reordering the DDS-index layout.
pad = self.torques.shape[0] - self.torque_limit.shape[0]
self.torque_limit = np.concatenate((self.torque_limit, np.full(pad, np.inf)))
if self.torque_limit.shape != self.torques.shape:
raise ValueError("motor_effort_limit_list must match the scene's actuator count")
base_body = self.mj_model.body("pelvis").id
self.world_object_joints = np.asarray(
[
i
for i in range(self.mj_model.njnt)
if self.mj_model.jnt_type[i] == mujoco.mjtJoint.mjJNT_FREE
and self.mj_model.jnt_bodyid[i] != base_body
],
dtype=int,
)
object_addresses = self.mj_model.jnt_qposadr[self.world_object_joints]
self.world_object_qpos0 = self.mj_model.qpos0[object_addresses[:, None] + np.arange(7)]
self.world_rng = np.random.default_rng()
self._init_grasp()
self.hand_open_state = self._settle_open_hands() if self.config.get("WORLD") else None
# Jittered objects keep their edge this far inside the table top (a box geom named table_top).
self.world_support = None
if mujoco.mj_name2id(self.mj_model, mujoco.mjtObj.mjOBJ_GEOM, "table_top") >= 0:
top = self.mj_model.geom("table_top")
self.world_support = (top.pos[:2].copy(), top.size[:2].copy())
for name in self.camera_configs:
if mujoco.mj_name2id(self.mj_model, mujoco.mjtObj.mjOBJ_CAMERA, name) < 0:
raise ValueError(f"Camera {name!r} does not exist in {self.config['ROBOT_SCENE']}")
self.reset()
def init_renderers(self):
# Initialize camera renderers
self.renderers = {}
for camera_name, camera_config in self.camera_configs.items():
renderer = mujoco.Renderer(
self.mj_model, height=camera_config["height"], width=camera_config["width"]
)
self.renderers[camera_name] = renderer
def start_image_publish_subprocess(self, start_method: str = "spawn", camera_port: int = 5555):
"""Start image publishing subprocess using ZMQ"""
# Use spawn method for better GIL isolation, or configured method
if len(self.camera_configs) == 0:
print(
"Warning: No camera configs provided, image publishing subprocess will not be started"
)
return
start_method = self.config.get("MP_START_METHOD", "spawn")
self.image_publish_process = ImagePublishProcess(
camera_configs=self.camera_configs,
image_dt=self.image_dt,
zmq_port=camera_port,
start_method=start_method,
verbose=self.config.get("verbose", False),
)
self.image_publish_process.start_process()
print(f"✓ Started image publishing subprocess on ZMQ port {camera_port}")
def compute_body_torques(self) -> np.ndarray:
"""Compute body torques based on the current robot state"""
body_torques = np.zeros(self.num_body_dof)
if self.unitree_bridge is not None and self.unitree_bridge.low_cmd:
# DDS command slots may be sparse (see UnitreeSdk2Bridge.joint_slots); index i is
# the MuJoCo actuator order, joint_slots[i] is the matching DDS motor_cmd slot.
for i, slot in enumerate(self.unitree_bridge.joint_slots):
if self.unitree_bridge.use_sensor:
body_torques[i] = (
self.unitree_bridge.low_cmd.motor_cmd[slot].tau
+ self.unitree_bridge.low_cmd.motor_cmd[slot].kp
* (self.unitree_bridge.low_cmd.motor_cmd[slot].q - self.mj_data.sensordata[i])
+ self.unitree_bridge.low_cmd.motor_cmd[slot].kd
* (
self.unitree_bridge.low_cmd.motor_cmd[slot].dq
- self.mj_data.sensordata[i + self.unitree_bridge.num_body_motor]
)
)
else:
body_torques[i] = (
self.unitree_bridge.low_cmd.motor_cmd[slot].tau
+ self.unitree_bridge.low_cmd.motor_cmd[slot].kp
* (
self.unitree_bridge.low_cmd.motor_cmd[slot].q
- self.mj_data.qpos[self.body_qpos_index[i]]
)
+ self.unitree_bridge.low_cmd.motor_cmd[slot].kd
* (
self.unitree_bridge.low_cmd.motor_cmd[slot].dq
- self.mj_data.qvel[self.body_dof_index[i]]
)
)
return body_torques
def compute_hand_torques(self) -> np.ndarray:
"""Compute hand torques based on the current robot state"""
left_hand_torques = np.zeros(self.num_hand_dof)
right_hand_torques = np.zeros(self.num_hand_dof)
if self.unitree_bridge is not None and self.unitree_bridge.low_cmd:
for i in range(self.unitree_bridge.num_hand_motor):
left_hand_torques[i] = (
self.unitree_bridge.left_hand_cmd.motor_cmd[i].tau
+ self.unitree_bridge.left_hand_cmd.motor_cmd[i].kp
* (
self.unitree_bridge.left_hand_cmd.motor_cmd[i].q
- self.mj_data.qpos[self.left_hand_qpos_index[i]]
)
+ self.unitree_bridge.left_hand_cmd.motor_cmd[i].kd
* (
self.unitree_bridge.left_hand_cmd.motor_cmd[i].dq
- self.mj_data.qvel[self.left_hand_dof_index[i]]
)
)
right_hand_torques[i] = (
self.unitree_bridge.right_hand_cmd.motor_cmd[i].tau
+ self.unitree_bridge.right_hand_cmd.motor_cmd[i].kp
* (
self.unitree_bridge.right_hand_cmd.motor_cmd[i].q
- self.mj_data.qpos[self.right_hand_qpos_index[i]]
)
+ self.unitree_bridge.right_hand_cmd.motor_cmd[i].kd
* (
self.unitree_bridge.right_hand_cmd.motor_cmd[i].dq
- self.mj_data.qvel[self.right_hand_dof_index[i]]
)
)
return np.concatenate((left_hand_torques, right_hand_torques))
def compute_body_qpos(self) -> np.ndarray:
"""Compute body joint positions based on the current command"""
body_qpos = np.zeros(self.num_body_dof)
if self.unitree_bridge is not None and self.unitree_bridge.low_cmd:
for i, slot in enumerate(self.unitree_bridge.joint_slots):
body_qpos[i] = self.unitree_bridge.low_cmd.motor_cmd[slot].q
return body_qpos
def compute_hand_qpos(self) -> np.ndarray:
"""Compute hand joint positions based on the current command"""
hand_qpos = np.zeros(self.num_hand_dof * 2)
if self.unitree_bridge is not None and self.unitree_bridge.low_cmd:
for i in range(self.unitree_bridge.num_hand_motor):
hand_qpos[i] = self.unitree_bridge.left_hand_cmd.motor_cmd[i].q
hand_qpos[i + self.num_hand_dof] = self.unitree_bridge.right_hand_cmd.motor_cmd[i].q
return hand_qpos
def prepare_obs(self) -> Dict[str, any]:
"""Prepare observation dictionary from the current robot state"""
obs = {}
obs["floating_base_pose"] = self.mj_data.qpos[:7]
obs["floating_base_vel"] = self.mj_data.qvel[:6]
obs["floating_base_acc"] = self.mj_data.qacc[:6]
obs["secondary_imu_quat"] = self.mj_data.xquat[self.torso_index]
obs["secondary_imu_vel"] = self.mj_data.cvel[self.torso_index]
obs["body_q"] = self.mj_data.qpos[self.body_qpos_index]
obs["body_dq"] = self.mj_data.qvel[self.body_dof_index]
obs["body_ddq"] = self.mj_data.qacc[self.body_dof_index]
obs["body_tau_est"] = self.mj_data.actuator_force[self.body_actuator_index]
if self.num_hand_dof > 0:
obs["left_hand_q"] = self.mj_data.qpos[self.left_hand_qpos_index]
obs["left_hand_dq"] = self.mj_data.qvel[self.left_hand_dof_index]
obs["left_hand_ddq"] = self.mj_data.qacc[self.left_hand_dof_index]
obs["left_hand_tau_est"] = self.mj_data.actuator_force[self.left_hand_actuator_index]
obs["right_hand_q"] = self.mj_data.qpos[self.right_hand_qpos_index]
obs["right_hand_dq"] = self.mj_data.qvel[self.right_hand_dof_index]
obs["right_hand_ddq"] = self.mj_data.qacc[self.right_hand_dof_index]
obs["right_hand_tau_est"] = self.mj_data.actuator_force[self.right_hand_actuator_index]
obs["time"] = self.mj_data.time
return obs
def sim_step(self):
self.obs = self.prepare_obs()
self.unitree_bridge.PublishLowState(self.obs)
if self.unitree_bridge.joystick:
self.unitree_bridge.PublishWirelessController()
if self.config["ENABLE_ELASTIC_BAND"]:
if self.elastic_band.enable:
# Get Cartesian pose and velocity of the band_attached_link
pose = np.concatenate(
[
self.mj_data.xpos[self.band_attached_link], # link position in world
self.mj_data.xquat[
self.band_attached_link
], # link quaternion in world [w,x,y,z]
np.zeros(6), # placeholder for velocity
]
)
# Get velocity in world frame
mujoco.mj_objectVelocity(
self.mj_model,
self.mj_data,
mujoco.mjtObj.mjOBJ_BODY,
self.band_attached_link,
pose[7:13],
0, # 0 for world frame
)
# Reorder velocity from [ang, lin] to [lin, ang]
pose[7:10], pose[10:13] = pose[10:13], pose[7:10].copy()
self.mj_data.xfrc_applied[self.band_attached_link] = self.elastic_band.Advance(pose)
else:
# explicitly resetting the force when the band is not enabled
self.mj_data.xfrc_applied[self.band_attached_link] = np.zeros(6)
body_torques = self.compute_body_torques()
hand_torques = self.compute_hand_torques()
self.torques[self.body_actuator_index] = body_torques
if self.num_hand_dof > 0:
self.torques[self.left_hand_actuator_index] = hand_torques[: self.num_hand_dof]
self.torques[self.right_hand_actuator_index] = hand_torques[self.num_hand_dof :]
self.torques = np.clip(self.torques, -self.torque_limit, self.torque_limit)
# Only the DDS-managed slots (body + dex1/dex3 hands) are written here; any other
# actuator (e.g. the ZMQ-driven head/AmazingHand pan-tilt/fingers) keeps whatever
# ctrl another writer (SimHeadHandDevice) has set, instead of being reset to zero.
self.mj_data.ctrl[self.dds_actuator_index] = self.torques[self.dds_actuator_index]
mujoco.mj_step(self.mj_model, self.mj_data)
self._update_grasp()
# self.check_self_collision()
def kinematics_step(self):
"""
Run kinematics only: compute the qpos of the robot and directly set the qpos.
For debugging purposes.
"""
if self.unitree_bridge is not None:
self.unitree_bridge.PublishLowState(self.prepare_obs())
if self.unitree_bridge.joystick:
self.unitree_bridge.PublishWirelessController()
if self.config["ENABLE_ELASTIC_BAND"]:
if self.elastic_band.enable:
# Get Cartesian pose and velocity of the band_attached_link
pose = np.concatenate(
[
self.mj_data.xpos[self.band_attached_link], # link position in world
self.mj_data.xquat[
self.band_attached_link
], # link quaternion in world [w,x,y,z]
np.zeros(6), # placeholder for velocity
]
)
# Get velocity in world frame
mujoco.mj_objectVelocity(
self.mj_model,
self.mj_data,
mujoco.mjtObj.mjOBJ_BODY,
self.band_attached_link,
pose[7:13],
0, # 0 for world frame
)
# Reorder velocity from [ang, lin] to [lin, ang]
pose[7:10], pose[10:13] = pose[10:13], pose[7:10].copy()
self.mj_data.xfrc_applied[self.band_attached_link] = self.elastic_band.Advance(pose)
else:
# explicitly resetting the force when the band is not enabled
self.mj_data.xfrc_applied[self.band_attached_link] = np.zeros(6)
body_qpos = self.compute_body_qpos() # (num_body_dof,)
hand_qpos = self.compute_hand_qpos() # (num_hand_dof * 2,)
self.mj_data.qpos[self.body_qpos_index] = body_qpos
self.mj_data.qpos[self.left_hand_qpos_index] = hand_qpos[: self.num_hand_dof]
self.mj_data.qpos[self.right_hand_qpos_index] = hand_qpos[self.num_hand_dof :]
mujoco.mj_kinematics(self.mj_model, self.mj_data)
mujoco.mj_comPos(self.mj_model, self.mj_data)
def apply_perturbation(self, key):
"""Apply perturbation to the robot"""
# Add velocity perturbations in body frame
perturbation_x_body = 0.0 # forward/backward in body frame
perturbation_y_body = 0.0 # left/right in body frame
if key == "up":
perturbation_x_body = 1.0 # forward
elif key == "down":
perturbation_x_body = -1.0 # backward
elif key == "left":
perturbation_y_body = 1.0 # left
elif key == "right":
perturbation_y_body = -1.0 # right
# Transform body frame velocity to world frame using MuJoCo's rotation
vel_body = np.array([perturbation_x_body, perturbation_y_body, 0.0])
vel_world = np.zeros(3)
base_quat = self.mj_data.qpos[3:7] # [w, x, y, z] quaternion
# Use MuJoCo's robust quaternion rotation (handles invalid quaternions automatically)
mujoco.mju_rotVecQuat(vel_world, vel_body, base_quat)
# Apply to base linear velocity in world frame
self.mj_data.qvel[0] += vel_world[0] # world X velocity
self.mj_data.qvel[1] += vel_world[1] # world Y velocity
# Update dynamics after velocity change
mujoco.mj_forward(self.mj_model, self.mj_data)
def _viewer_key_callback(self, key):
"""Viewer keys drive the elastic band (7/8/9) and are forwarded to the keyboard teleop."""
self.elastic_band.MujuocoKeyCallback(key)
forward_key(key)
def update_viewer(self):
if self.viewer is not None:
self.viewer.sync()
def update_viewer_camera(self):
if self.viewer is not None:
if self.viewer.cam.type == mujoco.mjtCamera.mjCAMERA_TRACKING:
self.viewer.cam.type = mujoco.mjtCamera.mjCAMERA_FREE
else:
self.viewer.cam.type = mujoco.mjtCamera.mjCAMERA_TRACKING
def set_unitree_bridge(self, unitree_bridge):
"""Set the unitree bridge from the simulator"""
self.unitree_bridge = unitree_bridge
def get_privileged_obs(self):
"""Get privileged observation. Should be implemented by subclasses."""
return {}
def update_render_caches(self):
"""Update render cache and shared memory for subprocess."""
# Lazy init renderers on first call (creates OpenGL context in calling thread)
if not self._renderers_initialized and self.offscreen:
self.init_renderers()
self._renderers_initialized = True
print(f"✓ Renderers initialized lazily in thread {__import__('threading').current_thread().name}")
render_caches = {}
for camera_name, camera_config in self.camera_configs.items():
renderer = self.renderers.get(camera_name)
if renderer is None:
continue
if "params" in camera_config:
renderer.update_scene(self.mj_data, camera=camera_config["params"])
else:
renderer.update_scene(self.mj_data, camera=camera_name)
render_caches[camera_name + "_image"] = renderer.render()
# Update shared memory if image publishing process is available
if self.image_publish_process is not None:
self.image_publish_process.update_shared_memory(render_caches)
return render_caches
def handle_keyboard_button(self, key):
if self.elastic_band is not None:
self.elastic_band.handle_keyboard_button(key)
if key == "backspace":
self.reset()
if key == "v":
self.update_viewer_camera()
if key in ["up", "down", "left", "right"]:
self.apply_perturbation(key)
def check_fall(self):
"""Check if the robot has fallen"""
self.fall = False
if self.mj_data.qpos[2] < 0.2:
self.fall = True
print(f"Warning: Robot has fallen, height: {self.mj_data.qpos[2]:.3f} m")
if self.fall:
self.reset()
def check_self_collision(self):
"""Check for self-collision of the robot"""
robot_bodies = get_subtree_body_names(self.mj_model, self.mj_model.body(self.root_body).id)
self_collision, contact_bodies = check_contact(
self.mj_model, self.mj_data, robot_bodies, robot_bodies, return_all_contact_bodies=True
)
if self_collision:
print(f"Warning: Self-collision detected: {contact_bodies}")
return self_collision
def reset(self):
mujoco.mj_resetData(self.mj_model, self.mj_data)
# Set valid floating base quaternion (identity: w=1, x=y=z=0)
# mj_resetData sets qpos to zeros, which gives invalid [0,0,0,0] quaternion
self.mj_data.qpos[3:7] = [1.0, 0.0, 0.0, 0.0]
if self.config.get("END_EFFECTOR") == "dex1":
for side in ("left", "right"):
joints = getattr(self, side + "_hand_index")
addresses = getattr(self, side + "_hand_qpos_index")
self.mj_data.qpos[addresses] = self.mj_model.jnt_range[joints, 1]
self._clear_grasp()
self.reset_world_objects()
if self.hand_open_state is not None:
qpos_index, qpos, ctrl_index, ctrl = self.hand_open_state
self.mj_data.qpos[qpos_index] = qpos
self.mj_data.ctrl[ctrl_index] = ctrl
# Propagate qpos to derived quantities (xquat, xpos, etc.)
mujoco.mj_forward(self.mj_model, self.mj_data)
def _init_grasp(self):
"""Per-hand state of the world's kinematic grasp: the `{side}_grasp` weld, finger servos and tcp site."""
model = self.mj_model
self.grasp_hands = {}
self.grasp_state = {}
self.grasp_geoms = np.zeros(0, int)
for side, first in (("right", 1), ("left", 11)):
weld = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_EQUALITY, f"{side}_grasp")
servos = [f"{side}_hand_motor{first + i}_joint" for i in range(6)]
if weld < 0 or mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, servos[0]) < 0:
continue
low, high = (GRASP_ZONE[0].copy(), GRASP_ZONE[1].copy())
if side == "left":
low[1], high[1] = -GRASP_ZONE[1][1], -GRASP_ZONE[0][1]
self.grasp_hands[side] = {
"weld": weld,
"qpos": np.array([model.jnt_qposadr[model.joint(name).id] for name in servos]),
"site": model.site(f"{side}_tcp").id,
"zone": (low, high),
"object": int(model.eq_obj2id[weld]),
}
self.grasp_state[side] = False
if self.grasp_hands:
bodies = {hand["object"] for hand in self.grasp_hands.values()}
self.grasp_geoms = np.flatnonzero(np.isin(model.geom_bodyid, list(bodies)))
self.grasp_conaffinity = model.geom_conaffinity[self.grasp_geoms].copy()
def _clear_grasp(self):
"""Drop every kinematic grasp (reset): welds inactive, hand-object contacts back."""
for side in self.grasp_state:
self.grasp_state[side] = False
self.mj_data.eq_active[self.grasp_hands[side]["weld"]] = self.mj_model.eq_active0[self.grasp_hands[side]["weld"]]
if self.grasp_hands:
self.mj_model.geom_conaffinity[self.grasp_geoms] = self.grasp_conaffinity
def _update_grasp(self):
"""Attach the object to a hand that has closed on it, release it when the hand opens (WORLD_GRASP=attach)."""
model, data = self.mj_model, self.mj_data
for side, hand in self.grasp_hands.items():
closure = amazing_hand_closure(data.qpos[hand["qpos"]])
if self.grasp_state[side]:
if closure <= GRASP_RELEASE_CLOSURE:
self._release_grasp(side)
continue
if closure < GRASP_ATTACH_CLOSURE or any(self.grasp_state.values()):
continue
mujoco.mj_kinematics(model, data)
tcp = data.site_xpos[hand["site"]]
local = data.site_xmat[hand["site"]].reshape(3, 3).T @ (data.xpos[hand["object"]] - tcp)
if np.all(local >= hand["zone"][0]) and np.all(local <= hand["zone"][1]):
self._attach_grasp(side, closure)
def _attach_grasp(self, side, closure):
model, data, hand = self.mj_model, self.mj_data, self.grasp_hands[side]
weld = hand["weld"]
body1, body2 = model.eq_obj1id[weld], hand["object"]
relpos = np.zeros(3)
relquat = np.zeros(4)
inverse = np.zeros(4)
mujoco.mju_negQuat(inverse, data.xquat[body1])
mujoco.mju_rotVecQuat(relpos, data.xpos[body2] - data.xpos[body1], inverse)
mujoco.mju_mulQuat(relquat, inverse, data.xquat[body2])
model.eq_data[weld, :3] = 0.0
model.eq_data[weld, 3:6] = relpos
model.eq_data[weld, 6:10] = relquat
data.eq_active[weld] = 1
model.geom_conaffinity[self.grasp_geoms] = self.grasp_conaffinity & ~2
self.grasp_state[side] = True
logger.info("%s hand attached %s (closure %.2f)", side, model.body(body2).name, closure)
def _release_grasp(self, side):
hand = self.grasp_hands[side]
self.mj_data.eq_active[hand["weld"]] = 0
self.mj_model.geom_conaffinity[self.grasp_geoms] = self.grasp_conaffinity
self.grasp_state[side] = False
logger.info("%s hand released %s", side, self.mj_model.body(hand["object"]).name)
def _settle_open_hands(self):
"""The AmazingHand joint state of the open hand, so a world starts with open rather than curled hands.
Returns None without AmazingHands. The loop-closing passive joints need the hand to be simulated to
its open pose, with everything else of the robot held.
"""
model = self.mj_model
names = [f"{side}_hand_motor{first + i}_joint" for side, first in (("right", 1), ("left", 11)) for i in range(8)]
if mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_ACTUATOR, names[0]) < 0:
return None
ctrl_index = np.array([model.actuator(name).id for name in names])
ctrl = np.radians(np.tile(AMAZING_HAND_OPEN_DEG, 8))
hand_q = np.zeros(model.nq, bool)
hand_v = np.zeros(model.nv, bool)
for j in range(model.njnt):
if "_hand_" in model.joint(j).name:
size = (4, 3) if model.jnt_type[j] == mujoco.mjtJoint.mjJNT_BALL else (1, 1)
hand_q[model.jnt_qposadr[j] : model.jnt_qposadr[j] + size[0]] = True
hand_v[model.jnt_dofadr[j] : model.jnt_dofadr[j] + size[1]] = True
data = mujoco.MjData(model)
data.qpos[:] = model.qpos0
held = data.qpos.copy()
data.ctrl[ctrl_index] = ctrl
for _ in range(int(1.5 / model.opt.timestep)):
mujoco.mj_step(model, data)
data.qpos[~hand_q] = held[~hand_q]
data.qvel[~hand_v] = 0.0
qpos_index = np.flatnonzero(hand_q)
return qpos_index, data.qpos[qpos_index].copy(), ctrl_index, ctrl
def reset_world_objects(self):
"""Put every free object of the world back at its spawn pose, jittered by WORLD_RANDOMIZE.
mj_resetData already restores qpos0, but the pose is set explicitly so it never depends on that.
"""
jitter = float(self.config.get("WORLD_RANDOMIZE", 0.0) or 0.0)
for joint, spawn in zip(self.world_object_joints, self.world_object_qpos0):
address = self.mj_model.jnt_qposadr[joint]
dof = self.mj_model.jnt_dofadr[joint]
pose = spawn.copy()
if jitter > 0:
pose[:2] += self.world_rng.uniform(-jitter, jitter, 2)
if self.world_support is not None:
geom = np.flatnonzero(self.mj_model.geom_bodyid == self.mj_model.jnt_bodyid[joint])[0]
reach = self.mj_model.geom_size[geom, 0] + WORLD_EDGE_MARGIN
centre, half = self.world_support
pose[:2] = np.clip(pose[:2], centre - half + reach, centre + half - reach)
yaw = self.world_rng.uniform(-np.pi / 12, np.pi / 12)
mujoco.mju_mulQuat(pose[3:], np.array([np.cos(yaw / 2), 0.0, 0.0, np.sin(yaw / 2)]), spawn[3:])
self.mj_data.qpos[address : address + 7] = pose
self.mj_data.qvel[dof : dof + 6] = 0.0
class BaseSimulator:
"""Base simulator class that handles initialization and running of simulations"""
def __init__(self, config: Dict[str, any], env_name: str = "default", **kwargs):
config = select_end_effector(config)
self.config = config
self.env_name = env_name
# Initialize ROS 2 node (optional, only if rclpy is available)
if HAS_RCLPY:
if not rclpy.ok():
rclpy.init()
self.node = rclpy.create_node("sim_mujoco")
self.thread = threading.Thread(target=rclpy.spin, args=(self.node,), daemon=True)
self.thread.start()
else:
self.thread = None
executor = rclpy.get_global_executor()
self.node = executor.get_nodes()[0] # will only take the first node
else:
self.node = None
self.thread = None
# Set update frequencies
self.sim_dt = self.config["SIMULATE_DT"]
self.image_dt = self.config.get("IMAGE_DT", 0.033333)
self.viewer_dt = self.config.get("VIEWER_DT", 0.02)
# Create the environment
self.sim_env = DefaultEnv(config, env_name, **kwargs)
# Initialize the DDS communication layer - should be safe to call multiple times
try:
if self.config.get("INTERFACE", None):
ChannelFactoryInitialize(self.config["DOMAIN_ID"], self.config["INTERFACE"])
else:
ChannelFactoryInitialize(self.config["DOMAIN_ID"])
except Exception as e:
# If it fails because it's already initialized, that's okay
print(f"Note: Channel factory initialization attempt: {e}")
# Initialize the unitree bridge and pass it to the environment
self.init_unitree_bridge()
self.sim_env.set_unitree_bridge(self.unitree_bridge)
# Initialize additional components
self.init_subscriber()
self.init_publisher()
self.sim_thread = None
def start_as_thread(self):
# Create simulation thread
self.sim_thread = Thread(target=self.start)
self.sim_thread.start()
def start_image_publish_subprocess(self, start_method: str = "spawn", camera_port: int = 5555):
"""Start the image publish subprocess"""
self.sim_env.start_image_publish_subprocess(start_method, camera_port)
def init_subscriber(self):
"""Initialize subscribers. Can be overridden by subclasses."""
pass
def init_publisher(self):
"""Initialize publishers. Can be overridden by subclasses."""
pass
def init_unitree_bridge(self):
"""Initialize the unitree SDK bridge and auto-detect joystick."""
self.unitree_bridge = UnitreeSdk2Bridge(self.config)
self.unitree_bridge.SetupJoystick(
device_id=self.config.get("JOYSTICK_DEVICE", 0),
js_type=self.config.get("JOYSTICK_TYPE", "xbox"),
)
def start(self):
"""Main simulation loop"""
import time
sim_cnt = 0
last_time = time.time()
print(f"Starting simulation loop. Viewer: {self.sim_env.viewer is not None}")
try:
while (
self.sim_env.viewer and self.sim_env.viewer.is_running()
) or self.sim_env.viewer is None:
# Run simulation step
self.sim_env.sim_step()
# Update viewer at viewer rate
if sim_cnt % int(self.viewer_dt / self.sim_dt) == 0:
self.sim_env.update_viewer()
# Update render caches at image rate
if sim_cnt % int(self.image_dt / self.sim_dt) == 0:
self.sim_env.update_render_caches()
# Sleep to maintain correct rate (simple timing without ROS)
elapsed = time.time() - last_time
sleep_time = max(0, self.sim_dt - elapsed)
if sleep_time > 0:
time.sleep(sleep_time)
last_time = time.time()
sim_cnt += 1
print(f"Loop exited. Viewer running: {self.sim_env.viewer.is_running() if self.sim_env.viewer else 'No viewer'}")
except KeyboardInterrupt:
# User pressed Ctrl+C - exit cleanly
print("Keyboard interrupt received")
pass
except Exception as e:
print(f"Exception in simulation loop: {e}")
import traceback
traceback.print_exc()
self.close()
def __del__(self):
"""Clean up resources when simulator is deleted"""
self.close()
def reset(self):
"""Reset the simulation. Can be overridden by subclasses."""
self.unitree_bridge.reset()
self.sim_env.reset()
def close(self):
"""Close the simulation. Can be overridden by subclasses."""
try:
# Stop image publishing subprocess
if hasattr(self.sim_env, "image_publish_process") and self.sim_env.image_publish_process is not None:
self.sim_env.image_publish_process.stop()
self.sim_env.image_publish_process = None
# Close viewer
if hasattr(self.sim_env, "viewer") and self.sim_env.viewer is not None:
self.sim_env.viewer.close()
for renderer in self.sim_env.renderers.values():
renderer.close()
self.sim_env.renderers.clear()
self.sim_env._renderers_initialized = False
# Shutdown ROS (if available)
if HAS_RCLPY and rclpy.ok():
rclpy.shutdown()
except Exception as e:
print(f"Warning during close: {e}")
def get_privileged_obs(self):
obs = self.sim_env.get_privileged_obs()
# TODO: add ros2 topic to get privileged obs
return obs
def handle_keyboard_button(self, key):
# Only handles keyboard buttons for default env.
if self.env_name == "default":
self.sim_env.handle_keyboard_button(key)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Robot")
parser.add_argument(
"--config",
type=str,
default="./gr00t_wbc/control/main/teleop/configs/g1_29dof_gear_wbc.yaml",
help="config file",
)
args = parser.parse_args()
with open(args.config, "r") as file:
config = yaml.load(file, Loader=yaml.FullLoader)
if config.get("INTERFACE", None):
ChannelFactoryInitialize(config["DOMAIN_ID"], config["INTERFACE"])
else:
ChannelFactoryInitialize(config["DOMAIN_ID"])
simulation = BaseSimulator(config)
simulation.start_as_thread()