commit 631dfd0a2da6c01d2b9382675002ad1b293bef1a Author: cyy_mac Date: Sun Jun 21 19:45:04 2026 +0800 orin diff --git a/.DS_Store b/.DS_Store new file mode 100644 index 0000000..5725f1f Binary files /dev/null and b/.DS_Store differ diff --git a/deploy_isaaclab_onnx_no_torch_wtw_r2.py b/deploy_isaaclab_onnx_no_torch_wtw_r2.py new file mode 100644 index 0000000..191f583 --- /dev/null +++ b/deploy_isaaclab_onnx_no_torch_wtw_r2.py @@ -0,0 +1,541 @@ +#!/usr/bin/env python3 +""" +deploy_isaaclab_onnx_no_torch_wtw_r2.py + +No-torch ONNX deployment script for Unitree Go1 using WTW LCM bridge. + +R2 logic matches original Walk-These-Ways style: + startup: press R2 -> move to default pose -> press R2 -> start policy + runtime: press R2 -> return to default pose and pause -> press R2 -> resume policy + +It is NOT hold-R2-to-run. R2 is treated as a rising-edge event, using +StateEstimator.right_lower_right_switch_pressed by default. + +Logs model input/output and targets as JSONL when --log-dir is provided. +""" + +import argparse +import json +import signal +import time +from datetime import datetime +from pathlib import Path + +import lcm +import numpy as np +import onnxruntime as ort + +from go1_gym_deploy.lcm_types.pd_tau_targets_lcmt import pd_tau_targets_lcmt + +# ---------------- IsaacLab policy constants ---------------- +NUM_OBS = 48 +NUM_ACTIONS = 12 +OBS_SCALES = {"lin_vel": 2.0, "ang_vel": 0.5, "dof_pos": 1.0, "dof_vel": 0.05} +ACTION_SCALE = 0.25 +CLIP_OBSERVATIONS = 100.0 + +# IsaacLab order: +# [FL_hip, FR_hip, RL_hip, RR_hip, +# FL_thigh, FR_thigh, RL_thigh, RR_thigh, +# FL_calf, FR_calf, RL_calf, RR_calf] +DEFAULT_JOINT_ANGLES_ISAAC = np.array([ + 0.1, -0.1, 0.1, -0.1, + 0.8, 0.8, 1.0, 1.0, + -1.5, -1.5, -1.5, -1.5, +], dtype=np.float32) + +# WTW internal order: +# [FL_hip, FL_thigh, FL_calf, +# FR_hip, FR_thigh, FR_calf, +# RL_hip, RL_thigh, RL_calf, +# RR_hip, RR_thigh, RR_calf] +ISAAC_TO_WTW = np.array([0, 4, 8, 1, 5, 9, 2, 6, 10, 3, 7, 11], dtype=np.int64) +WTW_TO_ISAAC = np.array([0, 3, 6, 9, 1, 4, 7, 10, 2, 5, 8, 11], dtype=np.int64) + +# WTW internal -> Unitree/lcm_position order. Same idea as WTW joint_idxs. +WTW_TO_UNITREE = np.array([3, 4, 5, 0, 1, 2, 9, 10, 11, 6, 7, 8], dtype=np.int64) + +NAMES_ISAAC = [ + "FL_hip", "FR_hip", "RL_hip", "RR_hip", + "FL_thigh", "FR_thigh", "RL_thigh", "RR_thigh", + "FL_calf", "FR_calf", "RL_calf", "RR_calf", +] +NAMES_WTW = [ + "FL_hip", "FL_thigh", "FL_calf", + "FR_hip", "FR_thigh", "FR_calf", + "RL_hip", "RL_thigh", "RL_calf", + "RR_hip", "RR_thigh", "RR_calf", +] +NAMES_UNITREE = [ + "FR_hip", "FR_thigh", "FR_calf", + "FL_hip", "FL_thigh", "FL_calf", + "RR_hip", "RR_thigh", "RR_calf", + "RL_hip", "RL_thigh", "RL_calf", +] + +EXIT = False + +def _sig_handler(signum, frame): + global EXIT + EXIT = True + +signal.signal(signal.SIGINT, _sig_handler) +signal.signal(signal.SIGTERM, _sig_handler) + + +class FakeStateEstimator: + def __init__(self, auto_r2=False): + self.default_wtw = DEFAULT_JOINT_ANGLES_ISAAC[ISAAC_TO_WTW].copy() + self.right_lower_right_switch_pressed = bool(auto_r2) + self.right_lower_right_switch = int(auto_r2) + self.left_stick = [0.0, 0.0] + self.right_stick = [0.0, 0.0] + + def get_body_linear_vel(self): + return np.zeros(3, dtype=np.float32) + + def get_body_angular_vel(self): + return np.zeros(3, dtype=np.float32) + + def get_gravity_vector(self): + return np.array([0.0, 0.0, -1.0], dtype=np.float32) + + def get_dof_pos(self): + return self.default_wtw.copy() + + def get_dof_vel(self): + return np.zeros(12, dtype=np.float32) + + def get_command(self): + return np.zeros(19, dtype=np.float32) + + +class OnnxPolicy: + def __init__(self, onnx_path): + self.session = ort.InferenceSession(onnx_path, providers=["CPUExecutionProvider"]) + self.input_name = self.session.get_inputs()[0].name + print("[INFO] ONNX loaded:", onnx_path) + print("[INFO] Inputs :", [(i.name, i.shape, i.type) for i in self.session.get_inputs()]) + print("[INFO] Outputs:", [(o.name, o.shape, o.type) for o in self.session.get_outputs()]) + + def __call__(self, obs): + out = self.session.run(None, {self.input_name: obs.reshape(1, -1).astype(np.float32)})[0][0] + return np.asarray(out, dtype=np.float32) + + +class JsonlLogger: + def __init__(self, log_dir, args): + self.enabled = bool(log_dir) + self.fp = None + self.log_every = max(1, int(args.log_every)) + self.flush_every = max(1, int(args.log_flush_every)) + self.run_dir = None + if not self.enabled: + return + ts = datetime.now().strftime("%Y%m%d_%H%M%S") + self.run_dir = Path(log_dir).expanduser().resolve() / f"run_{ts}" + self.run_dir.mkdir(parents=True, exist_ok=True) + meta = { + "created_at": ts, + "args": vars(args), + "num_obs": NUM_OBS, + "num_actions": NUM_ACTIONS, + "obs_scales": OBS_SCALES, + "action_scale": ACTION_SCALE, + "default_joint_angles_isaac": DEFAULT_JOINT_ANGLES_ISAAC.tolist(), + "isaac_to_wtw": ISAAC_TO_WTW.tolist(), + "wtw_to_isaac": WTW_TO_ISAAC.tolist(), + "wtw_to_unitree": WTW_TO_UNITREE.tolist(), + "names_isaac": NAMES_ISAAC, + "names_wtw": NAMES_WTW, + "names_unitree": NAMES_UNITREE, + } + (self.run_dir / "metadata.json").write_text(json.dumps(meta, indent=2, ensure_ascii=False), encoding="utf-8") + self.fp = open(self.run_dir / "steps.jsonl", "a", encoding="utf-8") + print("[INFO] Runtime log dir:", self.run_dir) + + def arr(self, x): + return np.asarray(x, dtype=np.float32).reshape(-1).tolist() + + def log(self, step, **kw): + if not self.enabled or step % self.log_every != 0: + return + rec = {"step": int(step), "time_wall": time.time()} + for k, v in kw.items(): + if isinstance(v, np.ndarray): + rec[k] = self.arr(v) + elif isinstance(v, (np.float32, np.float64)): + rec[k] = float(v) + elif isinstance(v, (np.int32, np.int64)): + rec[k] = int(v) + else: + rec[k] = v + self.fp.write(json.dumps(rec, ensure_ascii=False) + "\n") + if step % self.flush_every == 0: + self.fp.flush() + + def close(self): + if self.fp: + self.fp.flush() + self.fp.close() + print("[INFO] Runtime log saved:", self.run_dir) + + +def safe_array(x, n, name): + arr = np.asarray(x, dtype=np.float32).reshape(-1) + if arr.shape[0] != n: + raise ValueError(f"{name} should have length {n}, got {arr.shape[0]}") + return np.nan_to_num(arr, nan=0.0, posinf=0.0, neginf=0.0) + + +def get_commands(args, se): + if args.use_rc: + if not args.real_state: + return np.zeros(3, dtype=np.float32) + cmd = np.asarray(se.get_command(), dtype=np.float32).reshape(-1) + cmd = np.nan_to_num(cmd, nan=0.0, posinf=0.0, neginf=0.0) + if cmd.shape[0] < 3: + raise RuntimeError(f"RC command should have at least 3 values, got {cmd.shape}") + return np.array([cmd[0] * args.rc_x_scale, cmd[1] * args.rc_y_scale, cmd[2] * args.rc_yaw_scale], dtype=np.float32) + return np.array([args.cmd_x, args.cmd_y, args.cmd_yaw], dtype=np.float32) + + +def compute_obs(se, commands, last_actions): + obs = np.zeros(NUM_OBS, dtype=np.float32) + base_lin_vel = safe_array(se.get_body_linear_vel(), 3, "body_linear_vel") + base_ang_vel = safe_array(se.get_body_angular_vel(), 3, "body_angular_vel") + gravity = safe_array(se.get_gravity_vector(), 3, "gravity_vector") + dof_pos_wtw = safe_array(se.get_dof_pos(), 12, "dof_pos") + dof_vel_wtw = safe_array(se.get_dof_vel(), 12, "dof_vel") + dof_pos_isaac = dof_pos_wtw[WTW_TO_ISAAC] + dof_vel_isaac = dof_vel_wtw[WTW_TO_ISAAC] + + obs[0:3] = base_lin_vel * OBS_SCALES["lin_vel"] + obs[3:6] = base_ang_vel * OBS_SCALES["ang_vel"] + obs[6:9] = gravity + obs[9:12] = commands * np.array([OBS_SCALES["lin_vel"], OBS_SCALES["lin_vel"], OBS_SCALES["ang_vel"]], dtype=np.float32) + obs[12:24] = (dof_pos_isaac - DEFAULT_JOINT_ANGLES_ISAAC) * OBS_SCALES["dof_pos"] + obs[24:36] = dof_vel_isaac * OBS_SCALES["dof_vel"] + obs[36:48] = last_actions + obs = np.clip(obs, -CLIP_OBSERVATIONS, CLIP_OBSERVATIONS) + obs = np.nan_to_num(obs, nan=0.0, posinf=0.0, neginf=0.0) + info = { + "base_lin_vel": base_lin_vel, + "base_ang_vel": base_ang_vel, + "projected_gravity": gravity, + "dof_pos_wtw": dof_pos_wtw, + "dof_vel_wtw": dof_vel_wtw, + "dof_pos_isaac": dof_pos_isaac, + "dof_vel_isaac": dof_vel_isaac, + } + return obs.astype(np.float32), info + + +def state_ok(info): + dof_pos = info["dof_pos_wtw"] + grav = info["projected_gravity"] + if not np.all(np.isfinite(dof_pos)): + return False, "dof_pos has NaN/Inf" + if not np.all(np.isfinite(grav)): + return False, "gravity has NaN/Inf" + if np.linalg.norm(dof_pos) < 1e-6: + return False, "dof_pos is all zeros; likely no real leg_control_data received" + gn = float(np.linalg.norm(grav)) + if gn < 0.5 or gn > 1.5: + return False, f"gravity norm suspicious: {gn:.3f}" + return True, "ok" + + +def action_to_targets(action, action_clip): + action_safe = np.clip(action, -action_clip, action_clip).astype(np.float32) + q_isaac = DEFAULT_JOINT_ANGLES_ISAAC + action_safe * ACTION_SCALE + q_wtw = q_isaac[ISAAC_TO_WTW] + q_unitree = q_wtw[WTW_TO_UNITREE] + return action_safe, q_isaac.astype(np.float32), q_wtw.astype(np.float32), q_unitree.astype(np.float32) + + +def default_targets(): + q_isaac = DEFAULT_JOINT_ANGLES_ISAAC.copy() + q_wtw = q_isaac[ISAAC_TO_WTW] + q_unitree = q_wtw[WTW_TO_UNITREE] + return q_isaac, q_wtw, q_unitree + + +def make_msg(q_unitree, args, msg_id): + msg = pd_tau_targets_lcmt() + msg.q_des = np.asarray(q_unitree, dtype=np.float64).reshape(12).tolist() + msg.qd_des = [0.0] * 12 + msg.tau_ff = [0.0] * 12 + msg.kp = [float(args.kp)] * 12 + msg.kd = [float(args.kd)] * 12 + msg.timestamp_us = int(time.time() * 1e6) + msg.id = int(msg_id) + msg.robot_id = 0 + msg.se_contactState = [0.0] * 4 + return msg + + +def publish(lc, q_unitree, args, msg_id): + lc.publish("pd_plustau_targets", make_msg(q_unitree, args, msg_id).encode()) + + +def consume_r2_event(args, se): + if args.fake_auto_r2 and not args.real_state: + return True + val = bool(getattr(se, args.r2_pressed_field, False)) + if val: + try: + setattr(se, args.r2_pressed_field, False) + except Exception: + pass + return val + + +def wait_r2(args, se, text): + print(text) + if args.fake_auto_r2 and not args.real_state: + print("[INFO] fake_auto_r2 enabled: continuing immediately.") + return + while not EXIT: + if consume_r2_event(args, se): + print("[INFO] R2 press detected.") + return + time.sleep(0.01) + + +def calibrate_default(args, lc, se, msg_id): + if not args.publish: + print("[INFO] calibration skipped because publish=False") + return msg_id + print("[INFO] Moving slowly to default pose...") + default_wtw = DEFAULT_JOINT_ANGLES_ISAAC[ISAAC_TO_WTW].copy() + current_wtw = safe_array(se.get_dof_pos(), 12, "dof_pos") + if args.real_state and not args.no_state_safety_check and np.linalg.norm(current_wtw) < 1e-6: + print("[WARN] current dof_pos is all zeros; skip calibration publish") + return msg_id + rel = current_wtw - default_wtw + for _ in range(300): + if EXIT: + break + if np.max(np.abs(rel)) <= 0.01: + break + rel -= np.clip(rel, -args.cal_step_rad, args.cal_step_rad) + q_wtw = default_wtw + rel + publish(lc, q_wtw[WTW_TO_UNITREE], args, msg_id) + msg_id += 1 + time.sleep(args.cal_dt) + _, _, q_default_unitree = default_targets() + for _ in range(10): + if EXIT: + break + publish(lc, q_default_unitree, args, msg_id) + msg_id += 1 + time.sleep(args.cal_dt) + print("[INFO] Default pose commanded.") + return msg_id + + +def hold_default_until_r2(args, lc, se, msg_id): + print("[WTW R2] Paused at default pose. Press R2 again to resume policy.") + _, _, q_default_unitree = default_targets() + dt = 1.0 / max(1.0, args.hold_default_rate_hz) + next_t = time.perf_counter() + while not EXIT: + if consume_r2_event(args, se): + print("[WTW R2] Resume press detected.") + return msg_id + if args.publish: + publish(lc, q_default_unitree, args, msg_id) + msg_id += 1 + next_t += dt + sleep = next_t - time.perf_counter() + if sleep > 0: + time.sleep(sleep) + else: + next_t = time.perf_counter() + return msg_id + + +def print_vec(title, names, values): + print(title) + for n, v in zip(names, values): + print(f" {n:10s}: {float(v): .4f}") + + +def print_debug(step, mode, args, commands, obs, info, action_raw, action_safe, q_isaac, q_wtw, q_unitree, ok, reason, sent, r2): + print("\n" + "=" * 80) + print(f"[STEP {step}] mode={mode}, publish={args.publish}, publish_sent={sent}, real_state={args.real_state}, use_rc={args.use_rc}, wtw_r2_logic={args.wtw_r2_logic}, r2_pressed={r2}") + print(f"[STATE CHECK] ok={ok}, reason={reason}") + print(f"[CMD] x={commands[0]:.3f}, y={commands[1]:.3f}, yaw={commands[2]:.3f}") + print("[OBS] lin:", np.round(info["base_lin_vel"], 4), "ang:", np.round(info["base_ang_vel"], 4), "grav:", np.round(info["projected_gravity"], 4)) + print("[OBS] obs[0:12]:", np.round(obs[0:12], 4)) + print("[ACTION raw ]", np.round(action_raw, 4)) + print("[ACTION safe]", np.round(action_safe, 4)) + print("[TARGET offset max rad]:", float(np.max(np.abs(action_safe * ACTION_SCALE)))) + print_vec("[TARGET IsaacLab order]", NAMES_ISAAC, q_isaac) + print_vec("[TARGET WTW internal order]", NAMES_WTW, q_wtw) + print_vec("[TARGET Unitree/lcm_position order]", NAMES_UNITREE, q_unitree) + print("=" * 80) + + +def create_se(args, lc): + if not args.real_state: + print("[INFO] Using FakeStateEstimator. No real robot state will be read.") + return FakeStateEstimator(auto_r2=args.fake_auto_r2) + from go1_gym_deploy.utils.cheetah_state_estimator import StateEstimator + print("[INFO] Using real WTW StateEstimator.") + se = StateEstimator(lc) + se.spin() + return se + + +def parse_args(): + p = argparse.ArgumentParser(description="No-torch ONNX Go1 deployment with original WTW R2 logic") + p.add_argument("--onnx", default="policy.onnx") + p.add_argument("--lcm-url", default="udpm://239.255.76.67:7667?ttl=255") + g = p.add_mutually_exclusive_group() + g.add_argument("--real-state", action="store_true") + g.add_argument("--fake-state", action="store_true") + p.add_argument("--publish", action="store_true") + p.add_argument("--allow-fake-publish", action="store_true") + p.add_argument("--no-state-safety-check", action="store_true") + p.add_argument("--rate-hz", type=float, default=50.0) + p.add_argument("--action-clip", type=float, default=0.3) + p.add_argument("--kp", type=float, default=20.0) + p.add_argument("--kd", type=float, default=0.5) + p.add_argument("--cmd-x", type=float, default=0.0) + p.add_argument("--cmd-y", type=float, default=0.0) + p.add_argument("--cmd-yaw", type=float, default=0.0) + p.add_argument("--use-rc", action="store_true") + p.add_argument("--rc-x-scale", type=float, default=0.5) + p.add_argument("--rc-y-scale", type=float, default=0.5) + p.add_argument("--rc-yaw-scale", type=float, default=0.5) + p.add_argument("--wtw-r2-logic", action="store_true", help="press R2 to calibrate/start; press R2 to pause/resume") + p.add_argument("--r2-pressed-field", default="right_lower_right_switch_pressed") + p.add_argument("--fake-auto-r2", action="store_true", help="auto-accept R2 waits in fake-state tests") + p.add_argument("--cal-step-rad", type=float, default=0.05) + p.add_argument("--cal-dt", type=float, default=0.05) + p.add_argument("--hold-default-rate-hz", type=float, default=20.0) + p.add_argument("--log-dir", default="") + p.add_argument("--log-every", type=int, default=1) + p.add_argument("--log-flush-every", type=int, default=10) + p.add_argument("--print-every", type=int, default=50) + p.add_argument("--warmup-steps", type=int, default=10) + p.add_argument("--max-steps", type=int, default=0) + return p.parse_args() + + +def main(): + args = parse_args() + if args.publish and not args.real_state and not args.allow_fake_publish: + raise RuntimeError("Refusing to publish with fake state. Add --allow-fake-publish only for offline tests.") + if args.rate_hz <= 0: + raise ValueError("--rate-hz must be positive") + + print("[INFO] Args:", vars(args)) + lc = lcm.LCM(args.lcm_url) + se = create_se(args, lc) + policy = OnnxPolicy(args.onnx) + logger = JsonlLogger(args.log_dir, args) + + last_actions = np.zeros(NUM_ACTIONS, dtype=np.float32) + step = 0 + msg_id = 0 + dt = 1.0 / args.rate_hz + next_t = time.perf_counter() + + if not args.publish: + print("[INFO] DRY-RUN: no pd_plustau_targets will be published.") + else: + print("[WARN] PUBLISH mode enabled. Hang up the robot for first tests.") + + try: + if args.wtw_r2_logic and args.publish: + wait_r2(args, se, "[WTW R2] About to calibrate; robot will stand. Press R2 to calibrate.") + msg_id = calibrate_default(args, lc, se, msg_id) + wait_r2(args, se, "[WTW R2] Starting pose calibrated. Press R2 to start ONNX controller.") + elif args.wtw_r2_logic and not args.publish: + print("[INFO] --wtw-r2-logic enabled but publish=False; startup R2 waiting/calibration skipped in dry-run.") + + print("[INFO] Starting loop. Ctrl+C to exit.") + while not EXIT: + t0 = time.perf_counter() + mode = "policy" + commands = get_commands(args, se) + obs, info = compute_obs(se, commands, last_actions) + ok, reason = (True, "ok") + if args.real_state and not args.no_state_safety_check: + ok, reason = state_ok(info) + + r2 = consume_r2_event(args, se) if args.wtw_r2_logic else False + if args.wtw_r2_logic and r2 and args.publish: + print("[WTW R2] R2 pressed during policy: pause and return to default pose.") + msg_id = calibrate_default(args, lc, se, msg_id) + msg_id = hold_default_until_r2(args, lc, se, msg_id) + last_actions[:] = 0.0 + step += 1 + continue + + action_raw = policy(obs) + action_safe, q_isaac, q_wtw, q_unitree = action_to_targets(action_raw, args.action_clip) + last_actions = action_safe.copy() + + publish_sent = False + should_publish = args.publish and step >= args.warmup_steps + if should_publish and args.real_state and not args.no_state_safety_check and not ok: + if step % max(1, args.print_every) == 0: + print(f"[WARN] Not publishing because state check failed: {reason}") + elif should_publish: + publish(lc, q_unitree, args, msg_id) + msg_id += 1 + publish_sent = True + elif args.publish and step < args.warmup_steps and step % max(1, args.print_every) == 0: + print(f"[INFO] Warmup step {step}/{args.warmup_steps}: not publishing yet.") + + loop_ms = (time.perf_counter() - t0) * 1000.0 + logger.log( + step, + mode=mode, + loop_ms=loop_ms, + commands_isaac=commands, + obs_isaac=obs, + action_raw_isaac=action_raw, + action_safe_isaac=action_safe, + joint_targets_isaac=q_isaac, + joint_targets_wtw=q_wtw, + joint_targets_unitree=q_unitree, + base_lin_vel=info["base_lin_vel"], + base_ang_vel=info["base_ang_vel"], + projected_gravity=info["projected_gravity"], + dof_pos_wtw=info["dof_pos_wtw"], + dof_vel_wtw=info["dof_vel_wtw"], + dof_pos_isaac=info["dof_pos_isaac"], + dof_vel_isaac=info["dof_vel_isaac"], + state_ok=ok, + state_reason=reason, + r2_pressed_event=r2, + publish_sent=publish_sent, + ) + + if step % max(1, args.print_every) == 0: + print_debug(step, mode, args, commands, obs, info, action_raw, action_safe, q_isaac, q_wtw, q_unitree, ok, reason, publish_sent, r2) + + step += 1 + if args.max_steps > 0 and step >= args.max_steps: + print("[INFO] max_steps reached. Exiting.") + break + + next_t += dt + sleep = next_t - time.perf_counter() + if sleep > 0: + time.sleep(sleep) + else: + next_t = time.perf_counter() + + finally: + logger.close() + print("[INFO] Exiting.") + + +if __name__ == "__main__": + main() diff --git a/go1_sim2sim.py b/go1_sim2sim.py new file mode 100644 index 0000000..d0ce6a7 --- /dev/null +++ b/go1_sim2sim.py @@ -0,0 +1,236 @@ +#!/usr/bin/env python3 +"""Go1 sim2sim MuJoCo viewer — Original 30k flat training. + +Usage: python go1_sim2sim.py + +Requires: mujoco, onnxruntime, pynput +Install: pip install mujoco onnxruntime pynput + +Controls: + W/S: forward/back Q/E: strafe left/right + A/D: rotate Space: stop R: reset Esc: quit +""" +import numpy as np, mujoco, onnxruntime as ort, os, time, threading, queue +from mujoco import viewer +from pynput import keyboard + +HERE = os.path.dirname(os.path.abspath(__file__)) +ONNX = os.path.join(HERE, "policy.onnx") + +# ── Parameters (original MotrixLab Go1 config) ── +NUM_OBS = 45 +KP, KD = 80.0, 0.5 # KD=0.5 + MuJoCo joint_damping(0.5) = 1.0 = training kd +ACTION_SCALE = 0.05 +CLIP = 23.7 +DEFAULT_ANGLES = np.array([ + -0.0, 0.9, -1.8, + 0.0, 0.9, -1.8, + -0.0, 0.9, -1.8, + 0.0, 0.9, -1.8, +], dtype=np.float32) + +# ── Keyboard ── +class KB: + def __init__(s): + s._q = queue.Queue(); s.running = True; s.held = set() + def _n(s, k): + try: + if hasattr(k, 'char') and k.char: return k.char.lower() + except: pass + return str(k).lower() + def _w(s): + while s.running: + try: + et, k = s._q.get(timeout=0.05) + n = s._n(k) + if et == 'press': s.held.add(n) + elif et == 'release': s.held.discard(n) + except queue.Empty: pass + def init(s): + s._l = keyboard.Listener( + on_press=lambda k: s._q.put(('press', k)), + on_release=lambda k: s._q.put(('release', k))) + s._l.start() + s._t = threading.Thread(target=s._w, daemon=True); s._t.start() + def keys(s): return s.held.copy() + def stop(s): s.running = False; s._l.stop() + +# ── Main ── +def main(): + # The model XML is embedded below + xml = ''' + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +''' + + model = mujoco.MjModel.from_xml_string(xml) + data = mujoco.MjData(model) + data.qpos[0:3] = [0, 0, 0.42] + data.qpos[3:7] = [1, 0, 0, 0] + data.qpos[7:19] = DEFAULT_ANGLES + mujoco.mj_forward(model, data) + + session = ort.InferenceSession(ONNX, providers=['CPUExecutionProvider']) + print(f"[Go1 sim2sim] ONNX={ONNX}") + print(f"[Go1 sim2sim] PD kp={KP} kd={KD+0.5} action_scale={ACTION_SCALE} obs={NUM_OBS}-dim") + print(f"[Go1 sim2sim] W/S前后 Q/E左右 A/D旋转 Space停 R重置 Esc退出") + + kb = KB(); kb.init() + view = viewer.launch_passive(model, data) + step, vx, vy, wz = 0, 0.0, 0.0, 0.0 + last_a = np.zeros(12, dtype=np.float32) + + while view.is_running(): + keys = kb.keys() + if 'escape' in keys: break + if 'r' in keys: + data.qpos[0:3] = [0, 0, 0.42] + data.qpos[3:7] = [1, 0, 0, 0] + data.qpos[7:19] = DEFAULT_ANGLES + data.qvel[:] = 0; last_a[:] = 0 + mujoco.mj_forward(model, data) + + vx = 1.0 if 'w' in keys else (-1.0 if 's' in keys else 0.0) + vy = 1.0 if 'q' in keys else (-1.0 if 'e' in keys else 0.0) + wz = 1.0 if 'a' in keys else (-1.0 if 'd' in keys else 0.0) + if ' ' in keys: vx = vy = wz = 0.0 + + if step % 2 == 0: # 100Hz control (MuJoCo dt=0.005) + obs = np.zeros(NUM_OBS, dtype=np.float32) + sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, "gyro") + adr = model.sensor_adr[sid] + obs[0:3] = data.sensordata[adr:adr+3] * 0.25 + R = data.xmat[1].reshape(3, 3) + obs[3:6] = (R.T @ np.array([0., 0., -1.])).astype(np.float32) + obs[6:18] = (data.qpos[7:19] - DEFAULT_ANGLES) * 1.0 + obs[18:30] = data.qvel[6:18] * 0.05 + obs[30:42] = last_a + obs[42:45] = np.array([vx, vy, wz]) * np.array([2., 2., 0.25]) + obs = np.clip(obs, -100., 100.) + action = session.run(None, {'observations': obs.reshape(1, -1).astype(np.float32)})[0][0] + action = np.clip(action, -CLIP, CLIP) + last_a = action.copy() + + target = DEFAULT_ANGLES + action * ACTION_SCALE + torques = KP * (target - data.qpos[7:19]) - KD * data.qvel[6:18] + data.ctrl[:] = np.clip(torques, -CLIP, CLIP) + mujoco.mj_step(model, data) + view.sync() + step += 1 + time.sleep(0.001) + + kb.stop(); view.close() + +if __name__ == "__main__": + main() diff --git a/policy.onnx b/policy.onnx new file mode 100644 index 0000000..e275514 Binary files /dev/null and b/policy.onnx differ