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}>(.*?){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