Files
Motrixlab/scripts/dreamwaq_sim2sim_mujoco.py

455 lines
16 KiB
Python

#!/usr/bin/env python3
"""Run and validate a DreamWaQ ONNX policy in MuJoCo.
By default the robot walks forward through the stairs_box course. The process
returns success only after it climbs the stairs, crosses the top platform,
descends to flat ground, and remains upright.
Controls:
W/S: forward/back Q/E: left/right A/D: rotate
Space: stop R: reset Esc: quit
"""
import argparse
import os
import queue
import signal
import sys
import threading
import time
import mujoco
import numpy as np
import onnxruntime as ort
from mujoco import viewer
g_exit_requested = False
signal.signal(signal.SIGINT, lambda *args: globals().update(g_exit_requested=True))
PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
XML_DIR = os.path.join(
PROJECT_DIR, "motrix_envs", "src", "motrix_envs", "locomotion", "go1", "xmls"
)
DEFAULT_ONNX = os.path.join(PROJECT_DIR, "exports_go1_dreamwaq", "policy.onnx")
NUM_OBS = 45
NUM_ACTIONS = 12
HISTORY_LEN = 5
ACTION_SCALE = 0.25
KP = 28.0
KD = 0.7
CLIP_ACTIONS = 100.0
CLIP_TORQUES = 80.0
CLIP_OBS = 100.0
MAX_VX, MAX_VY, MAX_WZ = 1.0, 1.0, 1.0
# MuJoCo actuator and qpos order: FR, FL, RR, RL; hip, thigh, calf.
DEFAULT_ANGLES = np.array(
[
-0.1, 0.8, -1.5,
0.1, 0.8, -1.5,
-0.1, 1.0, -1.5,
0.1, 1.0, -1.5,
],
dtype=np.float32,
)
class Keyboard:
def __init__(self):
self.events = queue.Queue()
self.running = True
self.held = set()
self.thread = None
self.listener = None
@staticmethod
def _name(key):
try:
if hasattr(key, "char") and key.char:
return key.char.lower()
except Exception:
pass
return str(key).lower()
def _worker(self):
while self.running:
try:
event_type, key = self.events.get(timeout=0.05)
name = self._name(key)
if event_type == "press":
self.held.add(name)
else:
self.held.discard(name)
except queue.Empty:
pass
def start(self):
from pynput import keyboard
self.listener = keyboard.Listener(
on_press=lambda key: self.events.put(("press", key)),
on_release=lambda key: self.events.put(("release", key)),
)
self.listener.start()
self.thread = threading.Thread(target=self._worker, daemon=True)
self.thread.start()
def stop(self):
self.running = False
if self.listener is not None:
self.listener.stop()
def get_sensor(model, data, name):
sensor_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, name)
if sensor_id < 0:
return None
address = model.sensor_adr[sensor_id]
dimension = model.sensor_dim[sensor_id]
return data.sensordata[address : address + dimension].copy()
def compute_obs(model, data, commands, last_action):
"""Build the 45-value observation in the same order used for training."""
obs = np.zeros(NUM_OBS, dtype=np.float32)
gyro = get_sensor(model, data, "gyro")
obs[0:3] = (gyro if gyro is not None else data.qvel[3:6]) * 0.25
gravity_world = model.opt.gravity.copy()
gravity_world /= np.linalg.norm(gravity_world)
rotation = data.xmat[1].reshape(3, 3)
obs[3:6] = (rotation.T @ gravity_world).astype(np.float32)
obs[6:9] = commands * np.array([2.0, 2.0, 0.25], dtype=np.float32)
obs[9:21] = data.qpos[7:19] - DEFAULT_ANGLES
obs[21:33] = data.qvel[6:18] * 0.05
obs[33:45] = last_action
return np.clip(obs, -CLIP_OBS, CLIP_OBS)
def body_pose(data, body_id):
rotation = data.xmat[body_id].reshape(3, 3)
roll = np.arctan2(rotation[2, 1], rotation[2, 2])
pitch = np.arcsin(np.clip(-rotation[2, 0], -1.0, 1.0))
yaw = np.arctan2(rotation[1, 0], rotation[0, 0])
return data.xpos[body_id].copy(), np.degrees([roll, pitch, yaw])
def geom_x_range(model, prefix):
ranges = []
for geom_id in range(model.ngeom):
name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_GEOM, geom_id) or ""
if name.startswith(prefix):
center = model.geom_pos[geom_id, 0]
half_size = model.geom_size[geom_id, 0]
ranges.append((center - half_size, center + half_size))
if not ranges:
return None
return min(start for start, _ in ranges), max(end for _, end in ranges)
def course_geometry(model):
up_range = geom_x_range(model, "step_up_")
down_range = geom_x_range(model, "step_down_")
platform_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_GEOM, "platform")
if up_range is None or down_range is None or platform_id < 0:
return None
platform_x = model.geom_pos[platform_id, 0]
platform_half = model.geom_size[platform_id, 0]
platform_top = model.geom_pos[platform_id, 2] + model.geom_size[platform_id, 2]
return {
"up_start": up_range[0],
"platform_start": platform_x - platform_half,
"platform_end": platform_x + platform_half,
"down_end": down_range[1],
"platform_top": platform_top,
"pass_x": down_range[1] + 0.6,
}
def stage_for_x(x, course):
if course is None:
return "terrain"
if x < course["up_start"]:
return "approach"
if x < course["platform_start"]:
return "ascending"
if x < course["platform_end"]:
return "platform"
if x < course["pass_x"]:
return "descending"
return "finish"
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--onnx", default=DEFAULT_ONNX)
parser.add_argument(
"--terrain",
default="stairs_box",
choices=[
"flat", "rough", "stairs", "dreamwaq", "stairs_test",
"stairs_box", "flat_stairs",
],
)
parser.add_argument("--level", type=int, default=0)
parser.add_argument(
"--forward-speed", type=float, default=0.5,
help="default autonomous forward command in m/s",
)
parser.add_argument("--timeout", type=float, default=30.0, help="validation timeout in simulation seconds")
parser.add_argument("--log-interval", type=float, default=0.5, help="pose log interval in simulation seconds")
parser.add_argument(
"--realtime-factor", type=float, default=1.0,
help="viewer playback speed relative to wall time (default: 1.0)",
)
parser.add_argument("--manual", action="store_true", help="start with zero velocity and use keyboard commands")
parser.add_argument("--headless", action="store_true", help="run validation without viewer or real-time delay")
parser.add_argument("--no-validation", action="store_true", help="do not stop with automatic PASS/FAIL")
return parser.parse_args()
def main():
args = parse_args()
if hasattr(sys.stdout, "reconfigure"):
sys.stdout.reconfigure(line_buffering=True)
if not os.path.isfile(args.onnx):
print(f"[ERROR] ONNX not found: {args.onnx}", file=sys.stderr)
return 2
if args.headless and args.no_validation:
print("[ERROR] --headless requires automatic validation", file=sys.stderr)
return 2
if args.realtime_factor <= 0.0:
print("[ERROR] --realtime-factor must be positive", file=sys.stderr)
return 2
terrain_map = {
"flat": "scene_dreamwaq_flat.xml",
"rough": "scene_rough_terrain.xml",
"stairs": "scene_stairs_terrain.xml",
"dreamwaq": "scene_dreamwaq_terrain.xml",
"stairs_test": "scene_stairs_test.xml",
"stairs_box": "scene_stairs_box.xml",
"flat_stairs": "scene_flat_stairs.xml",
}
xml_file = os.path.join(XML_DIR, terrain_map[args.terrain])
previous_cwd = os.getcwd()
os.chdir(XML_DIR)
try:
model = mujoco.MjModel.from_xml_path(xml_file)
finally:
os.chdir(previous_cwd)
data = mujoco.MjData(model)
if args.terrain == "flat_stairs":
level = max(0, min(1, args.level))
spawn_x, spawn_y = -12.0, 4.0 - level * 8.0
elif args.terrain == "stairs_test":
spawn_x, spawn_y = -7.5, -4.0
elif args.terrain == "stairs_box":
spawn_x, spawn_y = -2.0, 0.0
elif args.terrain == "dreamwaq":
level = max(0, min(9, args.level))
spawn_x = -8.0
spawn_y = 36.0 - level * 8.0
else:
spawn_x, spawn_y = 0.0, 0.0
def hfield_z(x, y):
floor_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_GEOM, "floor")
if floor_id < 0 or model.geom_type[floor_id] != mujoco.mjtGeom.mjGEOM_HFIELD:
return 0.0
hfield_id = model.geom_dataid[floor_id]
rows = int(model.hfield_nrow[hfield_id])
columns = int(model.hfield_ncol[hfield_id])
size_x, size_y, height, base = model.hfield_size[hfield_id]
address = model.hfield_adr[hfield_id]
samples = model.hfield_data[address : address + rows * columns].reshape(rows, columns)
geom_position = model.geom_pos[floor_id]
column = int(np.clip(((x - geom_position[0]) / size_x * 0.5 + 0.5) * (columns - 1), 0, columns - 1))
row = int(np.clip(((y - geom_position[1]) / size_y * 0.5 + 0.5) * (rows - 1), 0, rows - 1))
return float(geom_position[2] + base + samples[row, column] * height)
spawn_z = hfield_z(spawn_x, spawn_y) + 0.34
def reset_state():
mujoco.mj_resetData(model, data)
data.qpos[0:3] = [spawn_x, spawn_y, spawn_z]
data.qpos[3:7] = [1.0, 0.0, 0.0, 0.0]
data.qpos[7:19] = DEFAULT_ANGLES
mujoco.mj_forward(model, data)
reset_state()
session = ort.InferenceSession(args.onnx, providers=["CPUExecutionProvider"])
course = course_geometry(model) if args.terrain == "stairs_box" else None
validation_enabled = not args.no_validation
if validation_enabled and course is None:
print("[ERROR] automatic validation requires --terrain stairs_box", file=sys.stderr)
return 2
print(f"[DreamWaQ] {args.onnx}")
print(f"[Terrain] {args.terrain} spawn=({spawn_x:.2f}, {spawn_y:.2f}, {spawn_z:.2f})")
print(f"[Command] vx={0.0 if args.manual else args.forward_speed:.2f} m/s")
if course is not None:
print(
"[Course] up={up_start:.2f}->{platform_start:.2f}m "
"platform={platform_start:.2f}->{platform_end:.2f}m "
"down={platform_end:.2f}->{down_end:.2f}m pass_x={pass_x:.2f}m".format(**course)
)
if not args.headless:
print(f"[Viewer] realtime_factor={args.realtime_factor:.2f}x")
print("[CTRL] W/S forward/back, Q/E lateral, A/D yaw, Space stop, R reset, Esc quit")
keyboard = None
view = None
if not args.headless:
keyboard = Keyboard()
keyboard.start()
view = viewer.launch_passive(model, data)
trunk_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "trunk")
if view is not None:
view.cam.lookat = data.body(trunk_id).xpos.copy()
view.cam.type = mujoco.mjtCamera.mjCAMERA_TRACKING
view.cam.trackbodyid = trunk_id
last_action = np.zeros(NUM_ACTIONS, dtype=np.float32)
action = np.zeros(NUM_ACTIONS, dtype=np.float32)
history = np.zeros((1, HISTORY_LEN, NUM_OBS), dtype=np.float32)
history_initialized = False
decimation = 4
physics_step = 0
reached_top = False
fallen_since = None
next_log_time = 0.0
status = 0
wall_start = time.monotonic()
try:
while not g_exit_requested and (view is None or view.is_running()):
keys = keyboard.held if keyboard is not None else set()
if "escape" in keys:
break
if "r" in keys:
reset_state()
last_action.fill(0.0)
history.fill(0.0)
history_initialized = False
reached_top = False
fallen_since = None
next_log_time = 0.0
physics_step = 0
wall_start = time.monotonic()
print("[Reset]")
forward = 0.0 if args.manual else args.forward_speed
vx = MAX_VX if "w" in keys else (-MAX_VX if "s" in keys else forward)
vy = MAX_VY if "q" in keys else (-MAX_VY if "e" in keys else 0.0)
wz = MAX_WZ if "a" in keys else (-MAX_WZ if "d" in keys else 0.0)
if " " in keys:
vx = vy = wz = 0.0
if physics_step % decimation == 0:
command = np.array([vx, vy, wz], dtype=np.float32)
obs = compute_obs(model, data, command, last_action)
if not history_initialized:
history[:, -1] = obs
history_initialized = True
action = session.run(
None,
{
"obs": obs.reshape(1, -1),
"obs_history": history.reshape(1, -1),
},
)[0][0]
if not np.all(np.isfinite(action)):
print("[FAIL] policy produced a non-finite action", file=sys.stderr)
status = 1
break
action = np.clip(action, -CLIP_ACTIONS, CLIP_ACTIONS)
last_action = action.copy()
if physics_step > 0:
history[:, :-1] = history[:, 1:]
history[:, -1] = obs
target = DEFAULT_ANGLES + action * ACTION_SCALE
actuator_joints = model.actuator_trnid[:, 0]
target = np.clip(
target,
model.jnt_range[actuator_joints, 0],
model.jnt_range[actuator_joints, 1],
)
torque = KP * (target - data.qpos[7:19]) - KD * data.qvel[6:18]
data.ctrl[:] = np.clip(torque, -CLIP_TORQUES, CLIP_TORQUES)
mujoco.mj_step(model, data)
physics_step += 1
position, rpy = body_pose(data, trunk_id)
stage = stage_for_x(position[0], course)
if course is not None and position[2] >= spawn_z + 0.35 and position[0] >= course["platform_start"] - 0.5:
reached_top = True
if data.time >= next_log_time:
velocity = data.qvel[0:3]
print(
f"[Pose] t={data.time:5.1f}s stage={stage:10s} "
f"pos=({position[0]:6.2f},{position[1]:6.2f},{position[2]:5.2f}) "
f"rpy=({rpy[0]:6.1f},{rpy[1]:6.1f},{rpy[2]:6.1f})deg "
f"vel=({velocity[0]:5.2f},{velocity[1]:5.2f},{velocity[2]:5.2f})"
)
next_log_time += args.log_interval
tilt = max(abs(rpy[0]), abs(rpy[1]))
fallen = position[2] < 0.18 or tilt > 65.0
if data.time > 1.0 and fallen:
fallen_since = data.time if fallen_since is None else fallen_since
else:
fallen_since = None
if validation_enabled:
upright_at_finish = position[2] < spawn_z + 0.25 and tilt < 35.0
if reached_top and position[0] >= course["pass_x"] and upright_at_finish:
print(
f"[PASS] climbed and descended stairs in {data.time:.2f}s; "
f"final_pos=({position[0]:.2f}, {position[1]:.2f}, {position[2]:.2f}), "
f"final_rpy=({rpy[0]:.1f}, {rpy[1]:.1f}, {rpy[2]:.1f})deg"
)
break
if fallen_since is not None and data.time - fallen_since >= 0.5:
print(
f"[FAIL] robot fell at stage={stage}, t={data.time:.2f}s, "
f"pos=({position[0]:.2f}, {position[1]:.2f}, {position[2]:.2f}), "
f"rpy=({rpy[0]:.1f}, {rpy[1]:.1f}, {rpy[2]:.1f})deg",
file=sys.stderr,
)
status = 1
break
if data.time >= args.timeout:
print(
f"[FAIL] timeout after {data.time:.1f}s at stage={stage}; "
f"reached_top={reached_top}, x={position[0]:.2f}m",
file=sys.stderr,
)
status = 1
break
if view is not None:
view.sync()
deadline = wall_start + data.time / args.realtime_factor
remaining = deadline - time.monotonic()
if remaining > 0:
time.sleep(remaining)
finally:
if keyboard is not None:
keyboard.stop()
if view is not None:
view.close()
return status
if __name__ == "__main__":
sys.exit(main())