From 94855db3edbb73a86b71d286e2e2b0f148b71aee Mon Sep 17 00:00:00 2001 From: cyy_mac Date: Sun, 26 Jul 2026 21:41:44 +0800 Subject: [PATCH] =?UTF-8?q?add=20fastcpp=E7=89=88=E6=9C=AC=EF=BC=8C?= =?UTF-8?q?=E6=9D=BF=E7=AB=AF=E6=8F=90=E9=80=9F300=E5=80=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- deploy_45dim_rl_gym/README.md | 38 + .../deploy_go1_rlgym_pro_sdk_fastcpp.py | 1159 ++++++++++++++++ .../deploy_go1_rlgym_pro_sdk_lab_fastcpp.py | 1160 +++++++++++++++++ 3 files changed, 2357 insertions(+) create mode 100644 deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_fastcpp.py create mode 100644 deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_lab_fastcpp.py diff --git a/deploy_45dim_rl_gym/README.md b/deploy_45dim_rl_gym/README.md index 73ce584..55acc77 100644 --- a/deploy_45dim_rl_gym/README.md +++ b/deploy_45dim_rl_gym/README.md @@ -5,6 +5,9 @@ - `policy_15k.onnx`:RoboGauge Go1 RL-Gym 的 ONNX 策略 - `deploy_go1_onnx_mujoco.py`:MuJoCo 调试脚本 - `deploy_go1_rlgym_pro_sdk.py`:真机 Go1 PRO 低层部署脚本 +- `deploy_go1_rlgym_pro_sdk_fastcpp.py`:真机 Go1 PRO 低层部署脚本,独立使用 C++ LowCmd 构包/加密后端 +- `deploy_go1_rlgym_pro_sdk_lab.py`:RobotLab 策略真机部署脚本 +- `deploy_go1_rlgym_pro_sdk_lab_fastcpp.py`:RobotLab 策略真机部署脚本,独立使用 C++ LowCmd 构包/加密后端 该 ONNX 策略使用 45 维单帧观测,以及 5 帧、共 225 维的历史输入,按观测项分组堆叠: @@ -14,6 +17,41 @@ 关节顺序为 `FR, FL, RR, RL`,与 `go1_pro_sdk` 的电机顺序一致。 +## Python / C++ 构包版本隔离 + +默认真机入口仍然走 Python SDK 构包和 Blowfish: + +```bash +mjpython deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk.py ... +mjpython deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_lab.py ... +``` + +C++ 低层构包测试使用独立入口,不会改动默认 Python 版本: + +```bash +mjpython deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_fastcpp.py ... +mjpython deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_lab_fastcpp.py ... +``` + +C++ 入口启动时会打印: + +```text +[INFO] LowCmd backend: cpp_fast_lowcmd +``` + +日志 `metadata.json` 里也会写入: + +```json +"lowcmd_backend": "cpp_fast_lowcmd" +``` + +如果 C++ 扩展没有编译,fastcpp 入口会在连接机器人前直接报错,不会回退到 Python 构包。板端编译命令: + +```bash +cd /root/go1_pro_sdk/fast_lowcmd_cpp +PYTHONPATH=/root/go1_pro_sdk python3 setup.py build_ext --inplace +``` + ## 安全说明 直接低层控电机是危险操作。 diff --git a/deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_fastcpp.py b/deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_fastcpp.py new file mode 100644 index 0000000..f5343d9 --- /dev/null +++ b/deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_fastcpp.py @@ -0,0 +1,1159 @@ +#!/usr/bin/env python3 +""" +Deploy the RoboGauge Go1 45-dim RL-Gym ONNX policy on Unitree Go1 PRO. + +This uses go1_pro_sdk direct MCU control, not LCM or the official Unitree SDK. + +Policy: + - policy_15k.onnx + - single-frame obs: 45 dims + - ONNX input: 5-frame history, 225 dims, stacked by observation terms + - 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 sys +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() +SDK_FAST_LOW_CMD = HERE.parents[1] / "go1_pro_sdk" / "fast_lowcmd_cpp" +if SDK_FAST_LOW_CMD.exists(): + sys.path.insert(0, str(SDK_FAST_LOW_CMD)) +try: + from fast_lowcmd import FastLowCmdBuilder +except ImportError as exc: + raise RuntimeError( + "C++ fast LowCmd backend is required for this entrypoint. " + "Build it first: cd /root/go1_pro_sdk/fast_lowcmd_cpp && " + "PYTHONPATH=/root/go1_pro_sdk python3 setup.py build_ext --inplace" + ) from exc + +DEFAULT_ONNX = HERE / "policy_30k.onnx" +LOWCMD_BACKEND = "cpp_fast_lowcmd" +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 = 5 +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([2.0, 2.0, 0.25], 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", +] +_FAST_LOW_CMD_BUILDER = None +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 225-dim 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"rlgym_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"], + "lowcmd_backend": LOWCMD_BACKEND, + } + 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 get_fast_lowcmd_builder(): + global _FAST_LOW_CMD_BUILDER + if _FAST_LOW_CMD_BUILDER is None: + _FAST_LOW_CMD_BUILDER = FastLowCmdBuilder() + return _FAST_LOW_CMD_BUILDER + + +def send_lowcmd_fast(client, cmd): + raw = get_fast_lowcmd_builder().build_encrypted_lowcmd(cmd) + client.send_raw(raw) + + +def send_damping(client): + raw = get_fast_lowcmd_builder().build_encrypted_damping() + client.send_raw(raw) + + +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, + ) + send_lowcmd_fast(client, 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, + ) + send_lowcmd_fast(client, 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, + ) + send_lowcmd_fast(client, 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(f"[INFO] LowCmd backend: {LOWCMD_BACKEND}") + 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 RL-Gym 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 225-dim 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 RL-Gym ONNX on Go1 PRO") + parser.add_argument("--onnx", default=str(DEFAULT_ONNX), help="Path to policy_15k.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=80.0) + parser.add_argument("--kd", type=float, default=1.0) + 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=1.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/deploy_go1_rlgym_pro_sdk_lab_fastcpp.py b/deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_lab_fastcpp.py new file mode 100644 index 0000000..a8283ec --- /dev/null +++ b/deploy_45dim_rl_gym/deploy_go1_rlgym_pro_sdk_lab_fastcpp.py @@ -0,0 +1,1160 @@ +#!/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 sys +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() +SDK_FAST_LOW_CMD = HERE.parents[1] / "go1_pro_sdk" / "fast_lowcmd_cpp" +if SDK_FAST_LOW_CMD.exists(): + sys.path.insert(0, str(SDK_FAST_LOW_CMD)) +try: + from fast_lowcmd import FastLowCmdBuilder +except ImportError as exc: + raise RuntimeError( + "C++ fast LowCmd backend is required for this entrypoint. " + "Build it first: cd /root/go1_pro_sdk/fast_lowcmd_cpp && " + "PYTHONPATH=/root/go1_pro_sdk python3 setup.py build_ext --inplace" + ) from exc + +DEFAULT_ONNX = HERE / "policy_robotlab_6500.onnx" +LOWCMD_BACKEND = "cpp_fast_lowcmd" +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", +] +_FAST_LOW_CMD_BUILDER = None +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"], + "lowcmd_backend": LOWCMD_BACKEND, + } + 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 get_fast_lowcmd_builder(): + global _FAST_LOW_CMD_BUILDER + if _FAST_LOW_CMD_BUILDER is None: + _FAST_LOW_CMD_BUILDER = FastLowCmdBuilder() + return _FAST_LOW_CMD_BUILDER + + +def send_lowcmd_fast(client, cmd): + raw = get_fast_lowcmd_builder().build_encrypted_lowcmd(cmd) + client.send_raw(raw) + + +def send_damping(client): + raw = get_fast_lowcmd_builder().build_encrypted_damping() + client.send_raw(raw) + + +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, + ) + send_lowcmd_fast(client, 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, + ) + send_lowcmd_fast(client, 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, + ) + send_lowcmd_fast(client, 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(f"[INFO] LowCmd backend: {LOWCMD_BACKEND}") + 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()