diff --git a/deploy_45dim_rl_gym/deploy_go1_onnx_mujoco_lab.py b/deploy_45dim_rl_gym/deploy_go1_onnx_mujoco_lab.py new file mode 100644 index 0000000..02188d0 --- /dev/null +++ b/deploy_45dim_rl_gym/deploy_go1_onnx_mujoco_lab.py @@ -0,0 +1,738 @@ +#!/usr/bin/env python3 +""" +Go1 RobotLab ONNX Policy - MuJoCo Simulation Deployment. + +Loads the RobotLab ONNX policy exported from RoboGauge and runs it in MuJoCo +with PD position control. The ONNX model uses a 10-frame history +(450-dim stacked-by-terms input, stateless). + +Usage: + conda activate free_dog_sdk + MUJOCO_GL=glfw mjpython deploy_go1_onnx_mujoco_lab.py + MUJOCO_GL=glfw mjpython deploy_go1_onnx_mujoco_lab.py --terrain terrains/stairs/stairs_6.xml + +Controls (matching RoboGauge keyboard convention): + ↑ / ↓ forward / backward + ← / → yaw left / right + , / . strafe left / right + K stop + R reset robot + Esc quit + +Architecture: + ONNX input: obs [1, 450] - 10 frames x 45 dims, stacked by TERMS + ONNX output: actions [1, 12] + Control: target_q = default_q + 0.25 * action, PD with Kp=28, Kd=0.7 +""" + +import argparse +import json +import os +import re +import signal +import sys +import time +from collections import deque +from datetime import datetime +from pathlib import Path + +import mujoco +import numpy as np +import onnxruntime as ort +from mujoco import viewer + +# ── path setup ── +SCRIPT_DIR = Path(__file__).resolve().parent +DEFAULT_ONNX = str(SCRIPT_DIR / "policy_robotlab_6500.onnx") +ROBOT_XML = str(SCRIPT_DIR / "go1.xml") +TERRAINS_DIR = SCRIPT_DIR / "terrains" + +# ── policy constants (RoboGauge Go1Config) ── +NUM_OBS = 45 +NUM_ACTIONS = 12 +HISTORY_LEN = 10 # RobotLab ONNX history frames +ONNX_INPUT_DIM = 450 # 45 x 10, stacked by terms + +ACTION_SCALE = 0.25 +KP = 28.0 +KD = 0.7 +CLIP_ACTIONS = 100.0 +CLIP_OBS = 100.0 + +ANG_VEL_SCALE = 0.25 +CMD_SCALE = np.array([1.0, 1.0, 1.0], dtype=np.float32) + +MAX_LIN_VEL_X = 1.0 +MAX_LIN_VEL_Y = 0.5 +MAX_ANG_VEL = 1.0 + +DEFAULT_DOF_POS = np.array([ + -0.1, 0.8, -1.5, # FR_hip, FR_thigh, FR_calf + 0.1, 0.8, -1.5, # FL_hip, FL_thigh, FL_calf + -0.1, 1.0, -1.5, # RR_hip, RR_thigh, RR_calf + 0.1, 1.0, -1.5, # RL_hip, RL_thigh, RL_calf +], dtype=np.float32) + +# ── exit flag ── +EXIT = False + +def _sig_handler(signum, frame): + global EXIT + EXIT = True + +signal.signal(signal.SIGINT, _sig_handler) +signal.signal(signal.SIGTERM, _sig_handler) + + +# ═══════════════════════════════════════════════════════════════ +# XML builder — merge terrain + robot into one MuJoCo scene +# ═══════════════════════════════════════════════════════════════ + +def _xml_inner(text, tag): + """Extract inner content of the first ... in text.""" + m = re.search(rf"<{tag}>(.*?)", text, re.DOTALL) + return m.group(1).strip() if m else "" + + +def build_scene_xml(robot_xml_path, terrain_xml_path=None): + """Merge robot XML with optional terrain XML into a single scene model. + + The robot XML (go1.xml) has base_link as the root body with no free joint. + We add a free joint and optionally merge terrain worldbody/asset elements. + """ + robot = Path(robot_xml_path).read_text() + + # 1) Add free joint to base_link + robot = robot.replace( + '', + '\n ' + ) + + if terrain_xml_path is None: + # No terrain — add a simple flat floor + floor = ( + '\n \n' + ) + robot = robot.replace( + ' into robot + terrain_assets = _xml_inner(terrain_text, "asset") + if terrain_assets: + # Insert before closing of robot (or before if no asset) + robot = robot.replace('', '\n' + terrain_assets + '\n ', 1) + + # 3) Merge terrain settings + terrain_visual = _xml_inner(terrain_text, "visual") + if terrain_visual: + robot = robot.replace('', '\n' + terrain_visual + '\n ', 1) + + # 4) Merge terrain worldbody elements (lights, geoms, etc.) before base_link + terrain_wb = _xml_inner(terrain_text, "worldbody") + if terrain_wb: + # Remove elements from terrain worldbody (we only want geoms/lights/cameras) + terrain_wb_no_bodies = re.sub(r'', '', terrain_wb, flags=re.DOTALL) + robot = robot.replace( + '= 0 + has_imu_quat = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, "Body_Quat") >= 0 + print(f"[INFO] IMU sensors: Body_Gyro={has_imu_gyro}, Body_Quat={has_imu_quat}") + + # ── Load ONNX ── + session = ort.InferenceSession(args.onnx, providers=['CPUExecutionProvider']) + inp = session.get_inputs()[0] + print(f"[INFO] ONNX input : {inp.name} {inp.shape}") + for o in session.get_outputs(): + print(f"[INFO] ONNX output: {o.name} {o.shape}") + + # ── Init state ── + data.qpos[0:3] = np.array([args.spawn_x, args.spawn_y, args.spawn_z], dtype=np.float64) + data.qpos[3:7] = np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float64) + data.qpos[7:19] = DEFAULT_DOF_POS.astype(np.float64) + data.qvel[:] = 0.0 + data.ctrl[:] = 0.0 + mujoco.mj_forward(model, data) + + # ── Control state ── + ctrl_dt = 0.02 # 50 Hz + sim_dt = model.opt.timestep + steps_per_inference = max(1, int(ctrl_dt / sim_dt)) + print(f"[INFO] control: {ctrl_dt}s ({1/ctrl_dt:.0f}Hz), " + f"sim steps per inference: {steps_per_inference}") + + last_action = np.zeros(NUM_ACTIONS, dtype=np.float32) + action_raw = np.zeros(NUM_ACTIONS, dtype=np.float32) + action = np.zeros(NUM_ACTIONS, dtype=np.float32) + target_pos = DEFAULT_DOF_POS.copy() + prev_target_pos = DEFAULT_DOF_POS.copy() + obs_builder = ObsBuilder() + step_count = 0 + inference_step = 0 + + if sample_mode: + cmd = np.array([ + args.sample_cmd_x, + args.sample_cmd_y, + args.sample_cmd_yaw, + ], dtype=np.float32) + cmd = np.clip(cmd, [-MAX_LIN_VEL_X, -MAX_LIN_VEL_Y, -MAX_ANG_VEL], + [MAX_LIN_VEL_X, MAX_LIN_VEL_Y, MAX_ANG_VEL]) + logger = SampleLogger(args.sample_log_dir, terrain_path, cmd, args) + samples = [] + max_sim_steps = int(args.sample_seconds / sim_dt) + print( + f"[INFO] sampling: seconds={args.sample_seconds:.1f} " + f"cmd=({cmd[0]:.2f},{cmd[1]:.2f},{cmd[2]:.2f}) " + f"spawn=({args.spawn_x:.2f},{args.spawn_y:.2f},{args.spawn_z:.2f})" + ) + + while step_count < max_sim_steps and not EXIT: + if inference_step == 0: + imu_quat = read_sensor(model, data, "Body_Quat", 4) + imu_ang_vel = read_sensor(model, data, "Body_Gyro", 3) + quat_wxyz = imu_quat if imu_quat is not None else data.qpos[3:7].copy() + if imu_ang_vel is not None: + base_ang_vel_body = imu_ang_vel + else: + world_ang_vel = data.qvel[3:6].copy() + base_ang_vel_body = quat_rotate_inverse(quat_wxyz, world_ang_vel) + + q = data.qpos[7:19].copy() + dq = data.qvel[6:18].copy() + obs_single = obs_builder.build_single_obs( + base_ang_vel_body, quat_wxyz, cmd, q, dq, last_action) + onnx_input = obs_builder.build_onnx_input(obs_single) + outputs = session.run(None, {'obs': onnx_input}) + action_raw = outputs[0][0].astype(np.float32) + action = np.clip(action_raw, -args.clip_actions, args.clip_actions) + last_action = action.copy() + + target_pos_raw = DEFAULT_DOF_POS + action * ACTION_SCALE + if args.max_target_step > 0.0: + delta = np.clip( + target_pos_raw - prev_target_pos, + -args.max_target_step, + args.max_target_step, + ) + target_pos = prev_target_pos + delta + else: + target_pos = target_pos_raw + prev_target_pos = target_pos.copy() + + current_pos = data.qpos[7:19] + current_vel = data.qvel[6:18] + torques = KP * (target_pos - current_pos) - KD * current_vel + torques = np.clip(torques, -33.5, 33.5) + data.ctrl[:] = torques.astype(np.float64) + + mujoco.mj_step(model, data) + + if inference_step == 0: + quat_wxyz = data.qpos[3:7].copy() + rpy_deg = quat_to_rpy_deg(quat_wxyz) + fallen = bool( + data.qpos[2] < 0.16 + or abs(rpy_deg[0]) > 60.0 + or abs(rpy_deg[1]) > 60.0 + or not np.all(np.isfinite(data.qpos)) + ) + rec = { + "step": int(step_count), + "control_step": int(step_count // steps_per_inference), + "time_sim": float(data.time), + "cmd": cmd.astype(float).tolist(), + "base_pos": data.qpos[0:3].astype(float).tolist(), + "base_quat": quat_wxyz.astype(float).tolist(), + "rpy_deg": rpy_deg.astype(float).tolist(), + "base_lin_vel": data.qvel[0:3].astype(float).tolist(), + "base_ang_vel": data.qvel[3:6].astype(float).tolist(), + "dof_pos": current_pos.astype(float).tolist(), + "dof_vel": current_vel.astype(float).tolist(), + "action_raw": action_raw.astype(float).tolist(), + "action": action.astype(float).tolist(), + "target_pos": target_pos.astype(float).tolist(), + "target_offset": (target_pos - DEFAULT_DOF_POS).astype(float).tolist(), + "torques": torques.astype(float).tolist(), + "fallen": fallen, + } + logger.write(rec) + samples.append(rec) + if rec["control_step"] % 50 == 0: + print( + f"[sample {rec['control_step']:04d}] " + f"t={rec['time_sim']:.2f} x={rec['base_pos'][0]:.2f} " + f"z={rec['base_pos'][2]:.2f} rpy={np.round(rpy_deg, 1)} " + f"act_max={np.max(np.abs(action)):.2f} " + f"target_off={np.max(np.abs(target_pos - DEFAULT_DOF_POS)):.2f}" + ) + if fallen: + print(f"[WARN] sample stopped: fallen at t={data.time:.2f}s") + break + + step_count += 1 + inference_step = (inference_step + 1) % steps_per_inference + + summary = summarize_samples(samples) + (logger.run_dir / "summary.json").write_text( + json.dumps(summary, indent=2, ensure_ascii=False)) + logger.close() + print("[INFO] sample summary:") + print(json.dumps(summary, indent=2, ensure_ascii=False)) + print(f"[INFO] sample saved: {logger.run_dir}") + return 0 + + # ── Keyboard ── + kb = KbReader() + kb.start() + + # ── Viewer ── + view = viewer.launch_passive(model, data) + # Track the robot body + body_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "base_link") + if body_id >= 0: + view.cam.type = mujoco.mjtCamera.mjCAMERA_TRACKING + view.cam.trackbodyid = body_id + view.cam.distance = 2.5 + view.cam.elevation = -20 + view.cam.azimuth = 60 + print("[INFO] viewer launched") + + loop_start = time.time() + + while view.is_running() and not EXIT: + held = kb.snapshot() + + # ── Quit ── + if 'key.esc' in held: + break + + # ── Command from keyboard ── + vx, vy, yaw = get_command(held) + cmd = np.array([vx, vy, yaw], dtype=np.float32) + + # ── Reset ── + if 'r' in held: + data.qpos[0:3] = np.array([0.0, 0.0, 0.34], dtype=np.float64) + data.qpos[3:7] = np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float64) + data.qpos[7:19] = DEFAULT_DOF_POS.astype(np.float64) + data.qvel[:] = 0.0 + data.ctrl[:] = 0.0 + last_action = np.zeros(NUM_ACTIONS, dtype=np.float32) + action_raw = np.zeros(NUM_ACTIONS, dtype=np.float32) + action = np.zeros(NUM_ACTIONS, dtype=np.float32) + target_pos = DEFAULT_DOF_POS.copy() + prev_target_pos = DEFAULT_DOF_POS.copy() + obs_builder.reset() + mujoco.mj_forward(model, data) + print("[RESET]") + + # ── Inference ── + if inference_step == 0: + # Match RoboGauge: policy obs uses XML IMU gyro and framequat sensors. + # Fall back to qpos/qvel only for XMLs without those sensors. + imu_quat = read_sensor(model, data, "Body_Quat", 4) + imu_ang_vel = read_sensor(model, data, "Body_Gyro", 3) + quat_wxyz = imu_quat if imu_quat is not None else data.qpos[3:7].copy() + if imu_ang_vel is not None: + base_ang_vel_body = imu_ang_vel + else: + world_ang_vel = data.qvel[3:6].copy() + base_ang_vel_body = quat_rotate_inverse(quat_wxyz, world_ang_vel) + + # Joint state (qpos[7:19] is FR,FL,RR,RL — matches policy order) + q = data.qpos[7:19].copy() + dq = data.qvel[6:18].copy() + + # Build obs + obs_single = obs_builder.build_single_obs( + base_ang_vel_body, quat_wxyz, cmd, q, dq, last_action) + onnx_input = obs_builder.build_onnx_input(obs_single) + + # Run ONNX + outputs = session.run(None, {'obs': onnx_input}) + action_raw = outputs[0][0].astype(np.float32) # [12] + action = np.clip(action_raw, -args.clip_actions, args.clip_actions) + last_action = action.copy() + + target_pos_raw = DEFAULT_DOF_POS + action * ACTION_SCALE + if args.max_target_step > 0.0: + delta = np.clip( + target_pos_raw - prev_target_pos, + -args.max_target_step, + args.max_target_step, + ) + target_pos = prev_target_pos + delta + else: + target_pos = target_pos_raw + prev_target_pos = target_pos.copy() + + # ── PD control ── + current_pos = data.qpos[7:19] + current_vel = data.qvel[6:18] + torques = KP * (target_pos - current_pos) - KD * current_vel + torques = np.clip(torques, -33.5, 33.5) + data.ctrl[:] = torques.astype(np.float64) + + mujoco.mj_step(model, data) + view.sync() + + # Real-time sync — use sim_dt because step_count increments every sim step + expected_time = step_count * sim_dt + elapsed = time.time() - loop_start + if 0 < expected_time - elapsed < ctrl_dt: + time.sleep(expected_time - elapsed) + + step_count += 1 + inference_step = (inference_step + 1) % steps_per_inference + + # Periodic status + if step_count % 200 == 0: + z = data.qpos[2] + lin_vel_abs = np.linalg.norm(data.qvel[0:3]) + print(f"[{step_count}] cmd=({vx:.1f},{vy:.1f},{yaw:.1f}) " + f"z={z:.3f} |v|={lin_vel_abs:.2f} " + f"act[0:4]={np.round(action[:4], 3)}") + + kb.stop() + view.close() + print("[INFO] done.") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_lab.py b/deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_lab.py new file mode 100644 index 0000000..7c07d35 --- /dev/null +++ b/deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_lab.py @@ -0,0 +1,1130 @@ +#!/usr/bin/env python3 +""" +Deploy the RoboGauge Go1 45-dim RobotLab ONNX policy on Unitree Go1 PRO. + +This uses go1_pro_sdk direct MCU control, not LCM or the official Unitree SDK. + +Policy: + - policy_robotlab_6500.onnx + - single-frame obs: 45 dims + - ONNX input: 10-frame history, 450 dims, stacked by observation terms + - command scale: [1.0, 1.0, 1.0] + - joint order: FR, FL, RR, RL, matching go1_pro_sdk motor order + +Safety-first workflow: + 1. MONITOR: --monitor, no motor command + 2. OBS-CHECK: --obs-check, no motor command + 3. INFER-CHECK: --infer-check, ONNX only, no motor command + 4. STATE MACHINE: + IDLE -> CALIBRATE -> HOLD -> OBS_TEST -> INFER_TEST -> RL + R2 advances one layer, L2 emergency-stops to IDLE. + RL output is only enabled when --enable-rl is passed. + +Before low-level control, kill sport processes on the Pi: + ssh pi@192.168.123.161 + sudo pkill -9 -f keep_sport_alive + sudo pkill -9 -f Legged_sport + sudo pkill -9 -f appTransit + +Initial tests should be done with the robot suspended. For extra-conservative +low-level checks, override the default with --power-factor 1. +""" + +import argparse +import json +import signal +import subprocess +import time +from collections import deque +from datetime import datetime +from enum import Enum +from pathlib import Path + +import numpy as np +import onnxruntime as ort + +from go1_pro_sdk import ( + MCUClient, LowCmd, MotorCmd, MotorMode, + apply_safety, PowerProtectViolation, JOINT_NAMES, +) + + +HERE = Path(__file__).parent.resolve() +DEFAULT_ONNX = HERE / "policy_robotlab_6500.onnx" +SPORT_KILL_CMD = ( + 'ssh pi@192.168.123.161 "sudo pkill -9 -f keep_sport_alive; ' + 'sudo pkill -9 -f Legged_sport; sudo pkill -9 -f appTransit"' +) + +NUM_OBS = 45 +NUM_ACTIONS = 12 +HISTORY_LEN = 10 +ONNX_INPUT_DIM = NUM_OBS * HISTORY_LEN + +ACTION_SCALE = 0.25 +CLIP_OBS = 100.0 +ANG_VEL_SCALE = 0.25 +DOF_VEL_SCALE = 0.05 +CMD_SCALE = np.array([1.0, 1.0, 1.0], dtype=np.float32) + +MAX_LIN_VEL_X = 1.0 +MAX_LIN_VEL_Y = 0.5 +MAX_ANG_VEL_YAW = 1.0 + +DEFAULT_DOF_POS = np.array([ + -0.1, 0.8, -1.5, # FR_hip, FR_thigh, FR_calf + 0.1, 0.8, -1.5, # FL_hip, FL_thigh, FL_calf + -0.1, 1.0, -1.5, # RR_hip, RR_thigh, RR_calf + 0.1, 1.0, -1.5, # RL_hip, RL_thigh, RL_calf +], dtype=np.float32) +EXPECTED_SDK_JOINT_NAMES = [ + "FR_0", "FR_1", "FR_2", + "FL_0", "FL_1", "FL_2", + "RR_0", "RR_1", "RR_2", + "RL_0", "RL_1", "RL_2", +] +POLICY_JOINT_NAMES = [ + "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 State(Enum): + IDLE = "IDLE" + CALIBRATE = "CALIBRATE" + HOLD = "HOLD" + OBS_TEST = "OBS_TEST" + INFER_TEST = "INFER_TEST" + RL = "RL" + FAULT = "FAULT" + + +def get_projected_gravity(quat_wxyz): + qw, qx, qy, qz = quat_wxyz + g = np.zeros(3, dtype=np.float32) + g[0] = 2.0 * (-qz * qx + qw * qy) + g[1] = -2.0 * (qz * qy + qw * qx) + g[2] = 1.0 - 2.0 * (qw * qw + qz * qz) + return g + + +def as_np(values, dtype=np.float32): + return np.asarray(values, dtype=dtype) + + +def motor_pos(state): + return np.array([state.motorState[i].q for i in range(NUM_ACTIONS)], dtype=np.float32) + + +def motor_vel(state): + return np.array([state.motorState[i].dq for i in range(NUM_ACTIONS)], dtype=np.float32) + + +def motor_tau(state): + return np.array([state.motorState[i].tauEst for i in range(NUM_ACTIONS)], dtype=np.float32) + + +def validate_joint_order(): + sdk_names = list(JOINT_NAMES) + if sdk_names != EXPECTED_SDK_JOINT_NAMES: + raise RuntimeError( + "go1_pro_sdk JOINT_NAMES order mismatch.\n" + f" expected: {EXPECTED_SDK_JOINT_NAMES}\n" + f" actual : {sdk_names}" + ) + print("[INFO] Joint order check passed: SDK and policy both use FR, FL, RR, RL.") + for i, (sdk_name, policy_name, q0) in enumerate( + zip(sdk_names, POLICY_JOINT_NAMES, DEFAULT_DOF_POS)): + print(f" [{i:02d}] {sdk_name:4s} -> {policy_name:8s} default={q0:+.3f}") + + +def apply_deadzone(value, deadzone): + if deadzone <= 0.0: + return float(value) + mag = abs(float(value)) + if mag <= deadzone: + return 0.0 + return float(np.sign(value) * (mag - deadzone) / max(1e-6, 1.0 - deadzone)) + + +def get_command(state, args): + if args.no_rc: + cmd = np.array([args.cmd_x, args.cmd_y, args.cmd_yaw], dtype=np.float32) + else: + r = state.remote + ly = apply_deadzone(r.ly, args.rc_deadzone) + lx = apply_deadzone(r.lx, args.rc_deadzone) + rx = apply_deadzone(r.rx, args.rc_deadzone) + if args.swap_vy_yaw: + cmd = np.array([ + ly * args.rc_vx_scale, + -rx * args.rc_vy_scale, + -lx * args.rc_wz_scale, + ], dtype=np.float32) + else: + cmd = np.array([ + ly * args.rc_vx_scale, + -lx * args.rc_vy_scale, + -rx * args.rc_wz_scale, + ], dtype=np.float32) + + limits = np.array([MAX_LIN_VEL_X, MAX_LIN_VEL_Y, MAX_ANG_VEL_YAW], dtype=np.float32) + return np.clip(cmd, -limits, limits) + + +class CommandFilter: + def __init__(self, args): + self.alpha = float(args.cmd_ema_alpha) + self.max_step = np.array([ + args.max_cmd_step_x, + args.max_cmd_step_y, + args.max_cmd_step_yaw, + ], dtype=np.float32) + self.prev = np.zeros(3, dtype=np.float32) + + def reset(self): + self.prev[:] = 0.0 + + def update(self, raw_cmd): + cmd = np.asarray(raw_cmd, dtype=np.float32) + if 0.0 < self.alpha < 1.0: + cmd = self.alpha * cmd + (1.0 - self.alpha) * self.prev + if np.any(self.max_step > 0.0): + limit = np.where(self.max_step > 0.0, self.max_step, np.inf) + cmd = self.prev + np.clip(cmd - self.prev, -limit, limit) + self.prev = cmd.astype(np.float32) + return self.prev.copy() + + +class ObsHistoryBuilder: + """Build RoboGauge 45-dim obs and 450-dim RobotLab term-stacked ONNX history.""" + + def __init__(self): + self.history = deque(maxlen=HISTORY_LEN) + + def reset(self): + self.history.clear() + + def build_single(self, state, cmd, last_action): + quat = as_np(state.imu.quaternion) + gyro = as_np(state.imu.gyroscope) + q = motor_pos(state) + dq = motor_vel(state) + + obs = np.zeros(NUM_OBS, dtype=np.float32) + obs[0:3] = gyro * ANG_VEL_SCALE + obs[3:6] = get_projected_gravity(quat) + obs[6:9] = cmd * CMD_SCALE + obs[9:21] = q - DEFAULT_DOF_POS + obs[21:33] = dq * DOF_VEL_SCALE + obs[33:45] = last_action + obs = np.clip(obs, -CLIP_OBS, CLIP_OBS) + return np.nan_to_num(obs, nan=0.0, posinf=0.0, neginf=0.0) + + def build_onnx_input(self, obs_single): + self.history.append(obs_single.copy()) + frames = list(self.history) + while len(frames) < HISTORY_LEN: + frames.insert(0, np.zeros(NUM_OBS, dtype=np.float32)) + + term_dims = [3, 3, 3, 12, 12, 12] + chunks = [] + offset = 0 + for dim in term_dims: + for frame in frames: + chunks.append(frame[offset:offset + dim]) + offset += dim + obs = np.concatenate(chunks, dtype=np.float32).reshape(1, ONNX_INPUT_DIM) + return np.nan_to_num(obs, nan=0.0, posinf=0.0, neginf=0.0) + + +class OnnxPolicy: + def __init__(self, onnx_path): + self.session = ort.InferenceSession(str(onnx_path), providers=["CPUExecutionProvider"]) + self.input_name = self.session.get_inputs()[0].name + self.input_shape = self.session.get_inputs()[0].shape + self.outputs = [(o.name, o.shape) for o in self.session.get_outputs()] + if self.input_shape[-1] != ONNX_INPUT_DIM: + raise ValueError(f"ONNX input shape {self.input_shape} does not match {ONNX_INPUT_DIM}") + print(f"[INFO] ONNX: {onnx_path}") + print(f"[INFO] Input : {self.input_name} {self.input_shape}") + print(f"[INFO] Output: {self.outputs}") + + def __call__(self, onnx_input): + outputs = self.session.run(None, {self.input_name: onnx_input.astype(np.float32)}) + return np.asarray(outputs[0][0], dtype=np.float32) + + +class RCEdgeDetector: + def __init__(self): + self._prev = set() + + def update(self, state): + current = set(state.remote.pressed) + rising = current - self._prev + falling = self._prev - current + self._prev = current + return rising, falling + + +class JsonlLogger: + def __init__(self, log_dir, args): + self.enabled = bool(log_dir) + self.fp = None + self.run_dir = None + self.flush_every = max(1, int(args.log_flush_every)) + if not self.enabled: + return + + ts = datetime.now().strftime("%Y%m%d_%H%M%S") + self.run_dir = Path(log_dir).expanduser().resolve() / f"robotlab_go1_deploy_{ts}" + self.run_dir.mkdir(parents=True, exist_ok=True) + meta = { + "created_at": ts, + "num_obs": NUM_OBS, + "history_len": HISTORY_LEN, + "onnx_input_dim": ONNX_INPUT_DIM, + "action_scale": ACTION_SCALE, + "default_dof_pos": DEFAULT_DOF_POS.tolist(), + "joint_names_sdk": list(JOINT_NAMES), + "joint_names_sdk_expected": EXPECTED_SDK_JOINT_NAMES, + "joint_names_policy": POLICY_JOINT_NAMES, + "joint_order_policy": ["FR", "FL", "RR", "RL"], + } + for k, v in vars(args).items(): + if isinstance(v, (str, int, float, bool, type(None))): + meta[k] = v + (self.run_dir / "metadata.json").write_text(json.dumps(meta, indent=2, ensure_ascii=False)) + self.fp = open(self.run_dir / "steps.jsonl", "a", encoding="utf-8", buffering=1) + print(f"[INFO] Log dir: {self.run_dir}") + + def log(self, step, **kw): + if not self.enabled: + return + rec = {"step": int(step), "time_wall": time.time()} + for k, v in kw.items(): + if isinstance(v, np.ndarray): + rec[k] = np.asarray(v, dtype=np.float32).reshape(-1).tolist() + 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(f"[INFO] Log saved: {self.run_dir}") + + +def fmt_rc(state): + r = state.remote + btns = ",".join(r.pressed) if r.pressed else "none" + return ( + f"lx={r.lx:+.2f} ly={r.ly:+.2f} rx={r.rx:+.2f} ry={r.ry:+.2f} " + f"L2={r.L2:.2f} btns={btns}" + ) + + +def send_damping(client): + client.send(LowCmd()) + + +def send_hold_cmd(client, state, args): + cmd = LowCmd() + for j in range(NUM_ACTIONS): + cmd.set_motor(j, MotorCmd( + mode=MotorMode.Servo, + q=float(DEFAULT_DOF_POS[j]), + dq=0.0, + tau=0.0, + Kp=args.kp, + Kd=args.kd, + )) + apply_safety( + cmd, + state, + power_factor=args.power_factor, + position_limit_on=True, + position_protect_limit=None, + ) + client.send(cmd) + + +def send_position_cmd(client, state, targets, args): + cmd = LowCmd() + for j in range(NUM_ACTIONS): + cmd.set_motor(j, MotorCmd( + mode=MotorMode.Servo, + q=float(targets[j]), + dq=0.0, + tau=0.0, + Kp=args.kp, + Kd=args.kd, + )) + pp_limit = args.position_protect_limit if args.position_protect_limit > 0 else None + apply_safety( + cmd, + state, + power_factor=args.power_factor, + position_limit_on=True, + position_protect_limit=pp_limit, + ) + client.send(cmd) + + +def ramp_to_default(client, args, state, logger=None, step_base=0): + print("[INFO] Ramping to default pose...") + current = motor_pos(state) + error = current - DEFAULT_DOF_POS + max_error = float(np.max(np.abs(error))) + print(f"[INFO] Current max default-pose error: {max_error:.3f} rad") + if max_error < 0.05: + print("[INFO] Already near default pose.") + return state + + ramp_steps = max(1, int(args.ramp_time * args.ramp_hz)) + dt = 1.0 / args.ramp_hz + log_every = max(1, int(args.ramp_hz / 5.0)) + next_t = time.perf_counter() + + for i in range(ramp_steps): + if EXIT: + return state + new_state = client.recv_latest() + if new_state is not None: + state = new_state + + ratio = float(i + 1) / float(ramp_steps) + target = current + ratio * (DEFAULT_DOF_POS - current) + cmd = LowCmd() + for j in range(NUM_ACTIONS): + cmd.set_motor(j, MotorCmd( + mode=MotorMode.Servo, + q=float(target[j]), + dq=0.0, + tau=0.0, + Kp=args.kp_cal, + Kd=args.kd_cal, + )) + apply_safety( + cmd, + state, + power_factor=args.power_factor, + position_limit_on=True, + position_protect_limit=None, + ) + client.send(cmd) + if logger is not None and (i % log_every == 0 or i == ramp_steps - 1): + log_state(logger, step_base * 100000 + i, "RAMP", state, target=target) + if i % max(1, ramp_steps // 4) == 0: + actual = motor_pos(state) + print( + f" ramp {i:4d}/{ramp_steps} " + f"target_err={np.max(np.abs(target - DEFAULT_DOF_POS)):.3f} " + f"actual_err={np.max(np.abs(actual - DEFAULT_DOF_POS)):.3f}" + ) + next_t += dt + sleep = next_t - time.perf_counter() + if sleep > 0: + time.sleep(sleep) + else: + next_t = time.perf_counter() + + hold_steps = max(1, int(0.5 * args.ramp_hz)) + for i in range(hold_steps): + if EXIT: + return state + new_state = client.recv_latest() + if new_state is not None: + state = new_state + send_hold_cmd(client, state, args) + if logger is not None and (i % log_every == 0 or i == hold_steps - 1): + log_state(logger, step_base * 100000 + ramp_steps + i, "RAMP_HOLD", state, target=DEFAULT_DOF_POS) + next_t += dt + sleep = next_t - time.perf_counter() + if sleep > 0: + time.sleep(sleep) + else: + next_t = time.perf_counter() + + print("[INFO] Default pose reached.") + return state + + +def state_ok(state, args): + q = motor_pos(state) + dq = motor_vel(state) + gyro = as_np(state.imu.gyroscope) + quat = as_np(state.imu.quaternion) + grav = get_projected_gravity(quat) + rpy_deg = np.degrees(as_np(state.imu.rpy)) + + checks = [ + (np.all(np.isfinite(q)), "joint position is non-finite"), + (np.all(np.isfinite(dq)), "joint velocity is non-finite"), + (np.all(np.isfinite(gyro)), "gyro is non-finite"), + (np.all(np.isfinite(quat)), "quaternion is non-finite"), + (0.5 <= np.linalg.norm(grav) <= 1.5, f"gravity norm={np.linalg.norm(grav):.3f}"), + (np.max(np.abs(dq)) <= args.max_dof_vel, f"max dof vel={np.max(np.abs(dq)):.2f}"), + (np.max(np.abs(gyro)) <= args.max_gyro, f"max gyro={np.max(np.abs(gyro)):.2f}"), + (abs(rpy_deg[0]) <= args.max_roll_deg, f"roll={rpy_deg[0]:.1f} deg"), + (abs(rpy_deg[1]) <= args.max_pitch_deg, f"pitch={rpy_deg[1]:.1f} deg"), + ] + for ok, reason in checks: + if not ok: + return False, reason + return True, "ok" + + +def action_ok(action_raw, args): + if not np.all(np.isfinite(action_raw)): + return False, "action is non-finite" + max_abs = float(np.max(np.abs(action_raw))) + if max_abs > args.action_trip_limit: + return False, f"action abs {max_abs:.2f} > trip {args.action_trip_limit:.2f}" + return True, "ok" + + +def clamp_action(action_raw, args): + return np.clip(action_raw, -args.action_clip, args.action_clip).astype(np.float32) + + +def smooth_action(action, prev_action, args): + alpha = float(args.action_ema_alpha) + if 0.0 < alpha < 1.0: + return (alpha * action + (1.0 - alpha) * prev_action).astype(np.float32) + return action.astype(np.float32) + + +def limit_target_step(target, prev_target, args): + limit = float(args.max_target_step) + if limit <= 0: + return target.astype(np.float32) + delta = np.clip(target - prev_target, -limit, limit) + return (prev_target + delta).astype(np.float32) + + +def kill_sport_processes(host, user): + cmds = [ + "sudo pkill -9 -f keep_sport_alive", + "sudo pkill -9 -f Legged_sport", + "sudo pkill -9 -f appTransit", + ] + ssh_target = f"{user}@{host}" if user else host + print(f"[INFO] Equivalent manual command: {SPORT_KILL_CMD}") + print(f"[INFO] Killing sport processes on {ssh_target}...") + try: + result = subprocess.run( + ["ssh", ssh_target, " && ".join(cmds)], + capture_output=True, + text=True, + timeout=15, + ) + if result.returncode == 0: + print("[INFO] Sport processes killed.") + return True + stderr = result.stderr.strip() + if "no process" in stderr.lower() or not stderr: + print("[INFO] No sport processes found.") + return True + print(f"[WARN] SSH returned {result.returncode}: {stderr}") + except FileNotFoundError: + print("[WARN] ssh command not found; kill sport processes manually.") + except subprocess.TimeoutExpired: + print("[WARN] SSH timed out; check Pi network.") + except Exception as exc: + print(f"[WARN] Failed to kill sport processes: {exc}") + return False + + +def connect_client(args): + validate_joint_order() + + if args.kill_sport: + kill_sport_processes(args.pi_host, args.pi_user) + + print("[INFO] Connecting to MCU...") + client = MCUClient() + print("[INFO] Waking MCU...") + client.wake_mcu(n_frames=50, dt=0.01) + state = client.recv_state(timeout=2.0) + if state is None: + client.close() + raise RuntimeError("No LowState received. Check robot network and sport processes.") + + print(f"[INFO] Connected. Battery={state.bms.SOC}%") + print(f"[INFO] RPY deg: {np.round(np.degrees(as_np(state.imu.rpy)), 1)}") + print("[INFO] Initial joint positions (rad):") + q = motor_pos(state) + for i, name in enumerate(JOINT_NAMES): + print(f" [{i:02d}] {name:4s}: q={q[i]:+7.3f}, default={DEFAULT_DOF_POS[i]:+7.3f}") + return client, state + + +def log_state(logger, step, mode, state, cmd=None, cmd_raw=None, obs_single=None, action_raw=None, + action_safe=None, target=None, state_reason="ok"): + logger.log( + step, + mode=mode, + battery_soc=state.bms.SOC, + rc_lx=state.remote.lx, + rc_ly=state.remote.ly, + rc_rx=state.remote.rx, + rc_ry=state.remote.ry, + rc_buttons=state.remote.pressed, + imu_rpy_deg=np.degrees(as_np(state.imu.rpy)), + imu_quat=as_np(state.imu.quaternion), + base_ang_vel=as_np(state.imu.gyroscope), + projected_gravity=get_projected_gravity(as_np(state.imu.quaternion)), + dof_pos=motor_pos(state), + dof_vel=motor_vel(state), + tau_est=motor_tau(state), + commands_raw=np.zeros(3, dtype=np.float32) if cmd_raw is None else cmd_raw, + commands=np.zeros(3, dtype=np.float32) if cmd is None else cmd, + obs_single=np.zeros(NUM_OBS, dtype=np.float32) if obs_single is None else obs_single, + action_raw=np.zeros(NUM_ACTIONS, dtype=np.float32) if action_raw is None else action_raw, + action_safe=np.zeros(NUM_ACTIONS, dtype=np.float32) if action_safe is None else action_safe, + joint_targets=np.zeros(NUM_ACTIONS, dtype=np.float32) if target is None else target, + state_reason=state_reason, + ) + + +def run_monitor(args): + print(MONITOR_BANNER) + input("Press Enter to start monitor...") + logger = JsonlLogger(args.log_dir, args) + client = None + try: + client, state = connect_client(args) + edge = RCEdgeDetector() + edge.update(state) + step = 0 + dt = 1.0 / args.rate_hz + next_t = time.perf_counter() + while not EXIT: + new_state = client.recv_latest() + if new_state is not None: + state = new_state + rising, falling = edge.update(state) + if step % args.print_every == 0: + rpy = np.degrees(as_np(state.imu.rpy)) + print(f"\n[MONITOR {step}] bat={state.bms.SOC}% rpy={np.round(rpy, 1)}") + print(f" RC: {fmt_rc(state)}") + if rising: + print(f" rising: {sorted(rising)}") + if falling: + print(f" falling: {sorted(falling)}") + print(f" q: {np.round(motor_pos(state), 3)}") + print(f" dq: {np.round(motor_vel(state), 3)}") + log_state(logger, step, "MONITOR", state) + step += 1 + if args.max_steps > 0 and step >= args.max_steps: + 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() + if client is not None: + client.close() + + +def run_obs_check(args): + print(OBS_BANNER) + input("Press Enter to start obs-check...") + logger = JsonlLogger(args.log_dir, args) + client = None + try: + client, state = connect_client(args) + obs_builder = ObsHistoryBuilder() + cmd_filter = CommandFilter(args) + last_action = np.zeros(NUM_ACTIONS, dtype=np.float32) + step = 0 + dt = 1.0 / args.rate_hz + next_t = time.perf_counter() + while not EXIT: + new_state = client.recv_latest() + if new_state is not None: + state = new_state + cmd_raw = get_command(state, args) + cmd = cmd_filter.update(cmd_raw) + obs_single = obs_builder.build_single(state, cmd, last_action) + onnx_input = obs_builder.build_onnx_input(obs_single) + if step % args.print_every == 0: + print(f"\n[OBS {step}] bat={state.bms.SOC}% {fmt_rc(state)}") + print(f" obs[0:3] gyro {np.round(obs_single[0:3], 4)}") + print(f" obs[3:6] gravity {np.round(obs_single[3:6], 4)}") + print(f" obs[6:9] cmd {np.round(obs_single[6:9], 4)} raw={np.round(cmd, 3)}") + print(f" obs[9:21] q-qd max={np.max(np.abs(obs_single[9:21])):.3f}") + print(f" obs[21:33] dq max={np.max(np.abs(obs_single[21:33])):.3f}") + print(f" onnx_input shape={onnx_input.shape} min={onnx_input.min():.3f} max={onnx_input.max():.3f}") + log_state(logger, step, "OBS_CHECK", state, cmd=cmd, cmd_raw=cmd_raw, obs_single=obs_single) + step += 1 + if args.max_steps > 0 and step >= args.max_steps: + 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() + if client is not None: + client.close() + + +def run_infer_check(args): + print(INFER_BANNER) + input("Press Enter to start infer-check...") + logger = JsonlLogger(args.log_dir, args) + client = None + try: + policy = OnnxPolicy(args.onnx) + client, state = connect_client(args) + obs_builder = ObsHistoryBuilder() + cmd_filter = CommandFilter(args) + last_action = np.zeros(NUM_ACTIONS, dtype=np.float32) + prev_action = np.zeros(NUM_ACTIONS, dtype=np.float32) + step = 0 + dt = 1.0 / args.rate_hz + next_t = time.perf_counter() + while not EXIT: + new_state = client.recv_latest() + if new_state is not None: + state = new_state + cmd_raw = get_command(state, args) + cmd = cmd_filter.update(cmd_raw) + obs_single = obs_builder.build_single(state, cmd, last_action) + onnx_input = obs_builder.build_onnx_input(obs_single) + action_raw = policy(onnx_input) + ok, reason = action_ok(action_raw, args) + action_safe = smooth_action(clamp_action(action_raw, args), prev_action, args) + prev_action = action_safe.copy() + last_action = action_safe.copy() + if step % args.print_every == 0: + print(f"\n[INFER {step}] ok={ok} reason={reason}") + print(f" cmd={np.round(cmd, 3)} action_raw={np.round(action_raw[:4], 3)} max={np.max(np.abs(action_raw)):.3f}") + print(f" action_safe max={np.max(np.abs(action_safe)):.3f}") + log_state( + logger, step, "INFER_CHECK", state, cmd=cmd, obs_single=obs_single, + cmd_raw=cmd_raw, + action_raw=action_raw, action_safe=action_safe, state_reason=reason, + ) + if not ok and args.trip_on_infer_check: + print(f"[FAULT] {reason}") + break + step += 1 + if args.max_steps > 0 and step >= args.max_steps: + 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() + if client is not None: + client.close() + + +STARTUP_BANNER = """ +============================================================ + RoboGauge Go1 RobotLab ONNX deployment + + State machine: + IDLE --R2--> CALIBRATE --> HOLD --R2--> OBS_TEST + OBS_TEST --R2--> INFER_TEST --R2--> RL + Any active state --L2--> IDLE damping + + RL motor commands require --enable-rl. Without it, R2 at INFER_TEST + will stay in INFER_TEST. + + Initial tests should be done with the robot suspended. + Use --kill-sport to run the Pi sport-process kill step before MCU control. + Manual equivalent: + ssh pi@192.168.123.161 "sudo pkill -9 -f keep_sport_alive; sudo pkill -9 -f Legged_sport; sudo pkill -9 -f appTransit" + After sport processes are killed, keep battery removal available as + the final stop method; the original sport-mode remote combo is not active. +============================================================ +""" + +MONITOR_BANNER = """ +============================================================ + MONITOR: read RC, IMU, and joint state only. No motor command. +============================================================ +""" + +OBS_BANNER = """ +============================================================ + OBS-CHECK: build 45-dim obs and 450-dim RobotLab history only. + No motor command. +============================================================ +""" + +INFER_BANNER = """ +============================================================ + INFER-CHECK: build obs and run ONNX only. + No motor command. +============================================================ +""" + + +def run_deploy(args): + print(STARTUP_BANNER) + input("Press Enter when ready...") + + policy = OnnxPolicy(args.onnx) + logger = JsonlLogger(args.log_dir, args) + client = None + state = None + sm_state = State.IDLE + step = 0 + cmd_raw = np.zeros(3, dtype=np.float32) + cmd = np.zeros(3, dtype=np.float32) + obs_single = None + action_raw = np.zeros(NUM_ACTIONS, dtype=np.float32) + action_safe = np.zeros(NUM_ACTIONS, dtype=np.float32) + target = DEFAULT_DOF_POS.copy() + reason = "ok" + + try: + client, state = connect_client(args) + edge = RCEdgeDetector() + edge.update(state) + obs_builder = ObsHistoryBuilder() + cmd_filter = CommandFilter(args) + + sm_state = State.IDLE + last_action = np.zeros(NUM_ACTIONS, dtype=np.float32) + prev_action = np.zeros(NUM_ACTIONS, dtype=np.float32) + prev_target = DEFAULT_DOF_POS.copy() + step = 0 + rl_step = 0 + dt = 1.0 / args.rate_hz + next_t = time.perf_counter() + + print("[INFO] R2 advances layers. L2 emergency-stops to IDLE.") + print("[INFO] Ctrl+C exits with safe_stop.") + + while not EXIT: + new_state = client.recv_latest() + if new_state is not None: + state = new_state + if state is None: + time.sleep(0.001) + continue + + rising, _ = edge.update(state) + r2_rose = "R2" in rising + l2_rose = "L2" in rising + + if l2_rose and sm_state != State.IDLE: + print(f"\n[L2] {sm_state.value} -> IDLE damping") + sm_state = State.IDLE + obs_builder.reset() + cmd_filter.reset() + last_action[:] = 0.0 + prev_action[:] = 0.0 + prev_target = DEFAULT_DOF_POS.copy() + rl_step = 0 + send_damping(client) + + cmd_raw = get_command(state, args) + if sm_state in (State.OBS_TEST, State.INFER_TEST, State.RL): + cmd = cmd_filter.update(cmd_raw) + else: + cmd_filter.reset() + cmd = cmd_raw + obs_single = None + action_raw = np.zeros(NUM_ACTIONS, dtype=np.float32) + action_safe = np.zeros(NUM_ACTIONS, dtype=np.float32) + target = DEFAULT_DOF_POS.copy() + reason = "ok" + + if sm_state == State.IDLE: + if step % 10 == 0: + send_damping(client) + if r2_rose: + print("\n[R2] IDLE -> CALIBRATE") + sm_state = State.CALIBRATE + state = ramp_to_default(client, args, state, logger=logger, step_base=step) + obs_builder.reset() + cmd_filter.reset() + last_action[:] = 0.0 + prev_action[:] = 0.0 + prev_target = DEFAULT_DOF_POS.copy() + sm_state = State.HOLD + print("[STATE] HOLD") + + elif sm_state == State.HOLD: + send_hold_cmd(client, state, args) + if r2_rose: + print("\n[R2] HOLD -> OBS_TEST") + obs_builder.reset() + cmd_filter.reset() + last_action[:] = 0.0 + prev_action[:] = 0.0 + sm_state = State.OBS_TEST + + elif sm_state == State.OBS_TEST: + send_hold_cmd(client, state, args) + obs_single = obs_builder.build_single(state, cmd, last_action) + obs_builder.build_onnx_input(obs_single) + ok, reason = state_ok(state, args) + if not ok: + print(f"\n[FAULT] OBS_TEST state check failed: {reason}") + sm_state = State.FAULT + elif r2_rose: + print("\n[R2] OBS_TEST -> INFER_TEST") + obs_builder.reset() + cmd_filter.reset() + last_action[:] = 0.0 + prev_action[:] = 0.0 + sm_state = State.INFER_TEST + + elif sm_state == State.INFER_TEST: + send_hold_cmd(client, state, args) + obs_single = obs_builder.build_single(state, cmd, last_action) + onnx_input = obs_builder.build_onnx_input(obs_single) + ok, reason = state_ok(state, args) + if ok: + action_raw = policy(onnx_input) + ok, reason = action_ok(action_raw, args) + if ok: + action_safe = smooth_action(clamp_action(action_raw, args), prev_action, args) + prev_action = action_safe.copy() + last_action = action_safe.copy() + else: + print(f"\n[FAULT] INFER_TEST failed: {reason}") + sm_state = State.FAULT + + if r2_rose and sm_state == State.INFER_TEST: + if not args.enable_rl: + print("\n[GUARD] RL blocked. Re-run with --enable-rl after OBS/INFER logs look safe.") + else: + print("\n[R2] INFER_TEST -> RL") + obs_builder.reset() + cmd_filter.reset() + last_action[:] = 0.0 + prev_action[:] = 0.0 + prev_target = DEFAULT_DOF_POS.copy() + rl_step = 0 + sm_state = State.RL + + elif sm_state == State.RL: + if r2_rose: + print("\n[R2] RL -> HOLD") + sm_state = State.HOLD + obs_builder.reset() + cmd_filter.reset() + last_action[:] = 0.0 + prev_action[:] = 0.0 + state = ramp_to_default(client, args, state, logger=logger, step_base=step) + prev_target = DEFAULT_DOF_POS.copy() + rl_step = 0 + continue + + obs_single = obs_builder.build_single(state, cmd, last_action) + onnx_input = obs_builder.build_onnx_input(obs_single) + ok, reason = state_ok(state, args) + if ok: + action_raw = policy(onnx_input) + ok, reason = action_ok(action_raw, args) + if not ok: + print(f"\n[FAULT] RL failed: {reason}") + sm_state = State.FAULT + send_damping(client) + else: + action_clipped = clamp_action(action_raw, args) + action_safe = smooth_action(action_clipped, prev_action, args) + target_raw = DEFAULT_DOF_POS + action_safe * ACTION_SCALE + target = limit_target_step(target_raw, prev_target, args) + prev_action = action_safe.copy() + last_action = action_safe.copy() + if rl_step >= args.warmup_steps: + prev_target = target.copy() + send_position_cmd(client, state, target, args) + else: + target = DEFAULT_DOF_POS.copy() + prev_target = DEFAULT_DOF_POS.copy() + send_hold_cmd(client, state, args) + rl_step += 1 + + elif sm_state == State.FAULT: + send_damping(client) + obs_builder.reset() + cmd_filter.reset() + last_action[:] = 0.0 + prev_action[:] = 0.0 + if r2_rose: + print("\n[R2] FAULT -> CALIBRATE") + sm_state = State.CALIBRATE + state = ramp_to_default(client, args, state, logger=logger, step_base=step) + prev_target = DEFAULT_DOF_POS.copy() + sm_state = State.HOLD + print("[STATE] HOLD") + + log_state( + logger, + step, + sm_state.value, + state, + cmd=cmd, + cmd_raw=cmd_raw, + obs_single=obs_single, + action_raw=action_raw, + action_safe=action_safe, + target=target, + state_reason=reason, + ) + + if step % args.print_every == 0: + rpy = np.degrees(as_np(state.imu.rpy)) + print( + f"\n[STEP {step}] state={sm_state.value} bat={state.bms.SOC}% " + f"rpy={np.round(rpy, 1)}" + ) + print(f" RC: {fmt_rc(state)}") + print(f" cmd={np.round(cmd, 3)} q={np.round(motor_pos(state), 2)}") + if sm_state in (State.INFER_TEST, State.RL, State.FAULT): + print( + f" action_raw_max={np.max(np.abs(action_raw)):.3f} " + f"action_safe_max={np.max(np.abs(action_safe)):.3f} reason={reason}" + ) + if sm_state == State.RL: + print(f" target={np.round(target, 2)} rl_step={rl_step}") + + step += 1 + if args.max_steps > 0 and step >= args.max_steps: + print("[INFO] max_steps reached.") + break + + next_t += dt + sleep = next_t - time.perf_counter() + if sleep > 0: + time.sleep(sleep) + else: + next_t = time.perf_counter() + + except PowerProtectViolation as exc: + reason = f"power protect: {exc}" + print(f"\n[FAULT] {reason}") + if client is not None: + send_damping(client) + if state is not None: + log_state( + logger, + step, + State.FAULT.value, + state, + cmd=cmd, + cmd_raw=cmd_raw, + obs_single=obs_single, + action_raw=action_raw, + action_safe=action_safe, + target=target, + state_reason=reason, + ) + + finally: + logger.close() + if client is not None: + print("[INFO] Safe stopping...") + client.safe_stop(n_frames=50, dt=0.002) + client.close() + print("[INFO] Done.") + + +def build_arg_parser(): + parser = argparse.ArgumentParser(description="Deploy RoboGauge Go1 RobotLab ONNX on Go1 PRO") + parser.add_argument("--onnx", default=str(DEFAULT_ONNX), help="Path to policy_robotlab_6500.onnx") + + parser.add_argument("--kill-sport", action="store_true", help="Kill Pi sport processes via SSH") + parser.add_argument("--pi-host", default="192.168.123.161") + parser.add_argument("--pi-user", default="pi") + + parser.add_argument("--monitor", action="store_true", help="Read RC/IMU/joints only; no motor command") + parser.add_argument("--obs-check", action="store_true", help="Build obs/history only; no motor command") + parser.add_argument("--infer-check", action="store_true", help="Run ONNX inference only; no motor command") + parser.add_argument("--enable-rl", action="store_true", help="Allow state machine to enter RL motor-control state") + parser.add_argument("--no-rc", action="store_true", help="Use fixed --cmd-* instead of RC sticks") + parser.add_argument("--swap-vy-yaw", action="store_true", help="Map left stick x to yaw and right stick x to vy") + + parser.add_argument("--kp", type=float, default=28.0) + parser.add_argument("--kd", type=float, default=0.7) + parser.add_argument("--kp-cal", type=float, default=20.0) + parser.add_argument("--kd-cal", type=float, default=1.0) + parser.add_argument("--power-factor", type=int, default=7) + parser.add_argument("--position-protect-limit", type=float, default=0.0) + + parser.add_argument("--rc-vx-scale", type=float, default=MAX_LIN_VEL_X) + parser.add_argument("--rc-vy-scale", type=float, default=MAX_LIN_VEL_Y) + parser.add_argument("--rc-wz-scale", type=float, default=MAX_ANG_VEL_YAW) + parser.add_argument("--rc-deadzone", type=float, default=0.05) + parser.add_argument("--cmd-ema-alpha", type=float, default=1.0) + parser.add_argument("--max-cmd-step-x", type=float, default=0.0) + parser.add_argument("--max-cmd-step-y", type=float, default=0.0) + parser.add_argument("--max-cmd-step-yaw", type=float, default=0.0) + parser.add_argument("--cmd-x", type=float, default=0.0) + parser.add_argument("--cmd-y", type=float, default=0.0) + parser.add_argument("--cmd-yaw", type=float, default=0.0) + + parser.add_argument("--rate-hz", type=float, default=50.0) + parser.add_argument("--ramp-time", type=float, default=5.0) + parser.add_argument("--ramp-hz", type=float, default=50.0) + parser.add_argument("--warmup-steps", type=int, default=50) + parser.add_argument("--max-steps", type=int, default=0) + parser.add_argument("--print-every", type=int, default=50) + + parser.add_argument("--max-roll-deg", type=float, default=35.0) + parser.add_argument("--max-pitch-deg", type=float, default=35.0) + parser.add_argument("--max-dof-vel", type=float, default=30.0) + parser.add_argument("--max-gyro", type=float, default=15.0) + parser.add_argument("--action-trip-limit", type=float, default=12.0) + parser.add_argument("--action-clip", type=float, default=4.0) + parser.add_argument("--action-ema-alpha", type=float, default=0.0) + parser.add_argument("--max-target-step", type=float, default=0.08) + parser.add_argument("--trip-on-infer-check", action="store_true") + + parser.add_argument("--log-dir", default=str(HERE / "logs"), help="JSONL log directory; empty disables logging") + parser.add_argument("--log-flush-every", type=int, default=50) + return parser + + +def main(): + args = build_arg_parser().parse_args() + if args.monitor: + return run_monitor(args) + if args.obs_check: + return run_obs_check(args) + if args.infer_check: + return run_infer_check(args) + return run_deploy(args) + + +if __name__ == "__main__": + main() diff --git a/deploy_45dim_rl_gym/policy_robotlab_6500.onnx b/deploy_45dim_rl_gym/policy_robotlab_6500.onnx new file mode 100644 index 0000000..935123a Binary files /dev/null and b/deploy_45dim_rl_gym/policy_robotlab_6500.onnx differ