#!/usr/bin/env python3 """ deploy_wtw_pro_sdk.py Walk-These-Ways (RMA) deployment on Unitree Go1 PRO via go1_pro_sdk. Direct MCU control — no LCM, no official Unitree SDK. Model: body_latest.jit (2102 → 12) + adaptation_module_latest.jit (2100 → 2) Observation: 70-dim, 30-step history (2100 dims total) Setup: cd /path/to/go1_pro_sdk && pip install -e . pip install torch onnxruntime numpy Before running: ssh pi@192.168.123.161 sudo pkill -9 -f keep_sport_alive; sudo pkill -9 -f Legged_sport; sudo pkill -9 -f appTransit Usage: python deploy_wtw_pro_sdk.py --kill-sport --log-dir ../logs """ import argparse import json import signal import subprocess import time from datetime import datetime from enum import Enum from pathlib import Path import numpy as np import torch from go1_pro_sdk import ( MCUClient, LowCmd, MotorCmd, MotorMode, apply_safety, JOINT_NAMES as SDK_JOINT_NAMES, ) HERE = Path(__file__).parent.resolve() # ─── WTW policy constants ─── NUM_OBS = 70 NUM_ACTIONS = 12 NUM_COMMANDS = 15 NUM_OBS_HISTORY = 30 # 30 steps of history OBS_HISTORY_DIM = NUM_OBS * NUM_OBS_HISTORY # 2100 BODY_INPUT_DIM = 2102 # 2100 (history) + 2 (latent) ADAPT_INPUT_DIM = 2100 LATENT_DIM = 2 # WTW joint order: FL→FR→RL→RR (per leg: hip/thigh/calf) # This is different from SDK order: FR→FL→RR→RL WTW_JOINT_NAMES = [ "FL_hip", "FL_thigh", "FL_calf", "FR_hip", "FR_thigh", "FR_calf", "RL_hip", "RL_thigh", "RL_calf", "RR_hip", "RR_thigh", "RR_calf", ] # SDK joint order: FR→FL→RR→RL (per leg: hip/thigh/calf) # Map: SDK index → WTW index SDK_TO_WTW = np.array([3, 4, 5, 0, 1, 2, 9, 10, 11, 6, 7, 8], dtype=np.int64) # Map: WTW index → SDK index WTW_TO_SDK = np.array([3, 4, 5, 0, 1, 2, 9, 10, 11, 6, 7, 8], dtype=np.int64) # Default joint angles in WTW order DEFAULT_ANGLES_WTW = np.array([ 0.1, 0.8, -1.5, # FL -0.1, 0.8, -1.5, # FR 0.1, 1.0, -1.5, # RL -0.1, 1.0, -1.5, # RR ], dtype=np.float32) # Default joint angles in SDK order DEFAULT_ANGLES_SDK = DEFAULT_ANGLES_WTW[WTW_TO_SDK] # Observation scales (WTW standard) OBS_SCALES = { "lin_vel": 2.0, "ang_vel": 0.25, "dof_pos": 1.0, "dof_vel": 0.05, "body_height_cmd": 2.0, "footswing_height_cmd": 0.15, "body_pitch_cmd": 0.3, "body_roll_cmd": 0.3, "stance_width_cmd": 1.0, "stance_length_cmd": 1.0, "aux_reward_cmd": 1.0, } COMMANDS_SCALE = np.array([ OBS_SCALES["lin_vel"], OBS_SCALES["lin_vel"], OBS_SCALES["ang_vel"], OBS_SCALES["body_height_cmd"], 1.0, 1.0, 1.0, 1.0, 1.0, OBS_SCALES["footswing_height_cmd"], OBS_SCALES["body_pitch_cmd"], OBS_SCALES["body_roll_cmd"], OBS_SCALES["stance_width_cmd"], OBS_SCALES["stance_length_cmd"], OBS_SCALES["aux_reward_cmd"], ], dtype=np.float32)[:NUM_COMMANDS] ACTION_SCALE = 0.25 HIP_SCALE_REDUCTION = 0.5 # hip joints get half action CLIP_ACTIONS = 10.0 CLIP_OBS = 100.0 EXIT = False def _sig_handler(signum, frame): global EXIT EXIT = True signal.signal(signal.SIGINT, _sig_handler) signal.signal(signal.SIGTERM, _sig_handler) # ─── State machine ─── class State(Enum): IDLE = "IDLE" CALIBRATE = "CALIBRATE" HOLD = "HOLD" RL = "RL" # ─── Quaternion math ─── def quat_to_rot_matrix(q): w, x, y, z = q return np.array([ [1 - 2*y*y - 2*z*z, 2*x*y - 2*w*z, 2*x*z + 2*w*y], [ 2*x*y + 2*w*z, 1 - 2*x*x - 2*z*z, 2*y*z - 2*w*x], [ 2*x*z - 2*w*y, 2*y*z + 2*w*x, 1 - 2*x*x - 2*y*y], ], dtype=np.float32) def get_projected_gravity(quaternion): R = quat_to_rot_matrix(quaternion) return (R.T @ np.array([0., 0., -1.], dtype=np.float32)).astype(np.float32) # ─── Observation ─── def build_commands_default(): """Build default command vector for trotting gait.""" cmd = np.zeros(NUM_COMMANDS, dtype=np.float32) cmd[0:3] = [0.0, 0.0, 0.0] # vx, vy, wz cmd[3] = 0.0 # height command (zero = nominal) cmd[4] = 3.0 # frequency (Hz) cmd[5] = 0.5 # phase (trot = 0.5 offset) cmd[6] = 0.0 # offset cmd[7] = 0.0 # bound cmd[8] = 0.5 # duration (stance ratio) cmd[9] = 0.15 # swing_height (matches reference footswing_height_cmd) cmd[10] = 0.0 # body_pitch cmd[11] = 0.0 # body_roll cmd[12] = 0.25 # stance_width cmd[13] = 0.4 # stance_length cmd[14] = 0.0 # aux_reward return cmd class ClockState: """Track gait indices and compute clock_inputs (4-dim sin per foot).""" def __init__(self): self.gait_indices = 0.0 self.dt = 0.01 # 100Hz def step(self, commands, dt=None): if dt is not None: self.dt = dt freq = commands[4] phase = commands[5] offset = commands[6] bound = commands[7] if NUM_COMMANDS > 8 else 0.0 self.gait_indices = (self.gait_indices + self.dt * freq) % 1.0 foot_indices = [ self.gait_indices + phase + offset + bound, # FL self.gait_indices + offset, # FR self.gait_indices + bound, # RL self.gait_indices + phase, # RR ] clock = np.array([np.sin(2 * np.pi * fi) for fi in foot_indices], dtype=np.float32) return clock def reset(self): self.gait_indices = 0.0 def compute_obs_wtw(imu, motor_states, commands, actions, last_actions, clock_inputs): """Build 70-dim observation matching WTW LCM agent layout.""" obs = np.zeros(NUM_OBS, dtype=np.float32) # 1. projected_gravity (3) obs[0:3] = get_projected_gravity(imu.quaternion) # 2. commands * scale (15) offset = 3 obs[offset:offset+NUM_COMMANDS] = commands * COMMANDS_SCALE offset += NUM_COMMANDS # 3. dof_pos_rel in WTW order (12) dof_pos_sdk = np.array([motor_states[i].q for i in range(12)], dtype=np.float32) dof_pos_wtw = dof_pos_sdk[SDK_TO_WTW] obs[offset:offset+12] = (dof_pos_wtw - DEFAULT_ANGLES_WTW) * OBS_SCALES["dof_pos"] offset += 12 # 4. dof_vel in WTW order (12) dof_vel_sdk = np.array([motor_states[i].dq for i in range(12)], dtype=np.float32) dof_vel_wtw = dof_vel_sdk[SDK_TO_WTW] obs[offset:offset+12] = dof_vel_wtw * OBS_SCALES["dof_vel"] offset += 12 # 5. actions clipped (12) obs[offset:offset+12] = np.clip(actions, -CLIP_ACTIONS, CLIP_ACTIONS) offset += 12 # 6. last_actions (12) obs[offset:offset+12] = last_actions offset += 12 # 7. clock_inputs (4) obs[offset:offset+4] = clock_inputs obs = np.clip(obs, -CLIP_OBS, CLIP_OBS) obs = np.nan_to_num(obs, nan=0.0, posinf=0.0, neginf=0.0) return obs # ─── Model ─── class WTWPolicy: def __init__(self, body_path, adapt_path): self.body = torch.jit.load(str(body_path), map_location='cpu') self.adapt = torch.jit.load(str(adapt_path), map_location='cpu') self.body.eval() self.adapt.eval() self.obs_history = torch.zeros(1, OBS_HISTORY_DIM, dtype=torch.float) self.latent = torch.zeros(1, LATENT_DIM, dtype=torch.float) print(f"[INFO] WTW body: {body_path}") print(f"[INFO] WTW adapt: {adapt_path}") print(f"[INFO] History: {NUM_OBS} obs × {NUM_OBS_HISTORY} steps = {OBS_HISTORY_DIM}") def reset(self): self.obs_history.zero_() self.latent.zero_() def __call__(self, obs): obs_t = torch.from_numpy(obs.reshape(1, -1)).float() # Update history: shift left, append new obs self.obs_history = torch.cat( (self.obs_history[:, NUM_OBS:], obs_t), dim=-1) # Adaptation module: history → latent with torch.no_grad(): self.latent = self.adapt(self.obs_history) # Body: [history, latent] → action body_input = torch.cat((self.obs_history, self.latent), dim=-1) with torch.no_grad(): action = self.body(body_input) return action.numpy().flatten().astype(np.float32) # ─── Remote controller ─── def get_rc_commands(state, base_cmd, args): r = state.remote base_cmd[0] = r.ly * args.rc_vx_scale base_cmd[1] = -r.lx * args.rc_vy_scale base_cmd[2] = -r.rx * args.rc_wz_scale return base_cmd 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 # ─── Safety wrappers ─── def send_hold_cmd(client, state, args): cmd = LowCmd() for j in range(12): cmd.set_motor(j, MotorCmd( mode=MotorMode.Servo, q=float(DEFAULT_ANGLES_SDK[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_rl_cmd(client, state, action_wtw, args): """Convert WTW action to SDK targets and send.""" # Scale action → position offset in WTW order offset_wtw = action_wtw * ACTION_SCALE # Apply hip scale reduction for i in [0, 3, 6, 9]: # hip indices in WTW order offset_wtw[i] *= HIP_SCALE_REDUCTION # Target in WTW order targets_wtw = DEFAULT_ANGLES_WTW + offset_wtw # Convert to SDK order targets_sdk = targets_wtw[WTW_TO_SDK] cmd = LowCmd() for j in range(12): cmd.set_motor(j, MotorCmd( mode=MotorMode.Servo, q=float(targets_sdk[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 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}" 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 or "no process" in result.stderr.lower(): print("[INFO] Sport processes killed.") return True print(f"[WARN] SSH returned {result.returncode}: {result.stderr.strip()}") return False except Exception as e: print(f"[WARN] Failed: {e}") return False # ─── JSONL logger ─── class JsonlLogger: def __init__(self, log_dir, args): self.enabled = bool(log_dir) self.fp = None self.run_dir = None self.flush_every = 50 if not self.enabled: return ts = datetime.now().strftime("%Y%m%d_%H%M%S") self.run_dir = Path(log_dir).expanduser().resolve() / f"wtw_deploy_{ts}" self.run_dir.mkdir(parents=True, exist_ok=True) meta = { "created_at": ts, "num_obs": NUM_OBS, "num_actions": NUM_ACTIONS, "num_commands": NUM_COMMANDS, "num_obs_history": NUM_OBS_HISTORY, "action_scale": ACTION_SCALE, "default_angles_wtw": DEFAULT_ANGLES_WTW.tolist(), "default_angles_sdk": DEFAULT_ANGLES_SDK.tolist(), "joint_names_wtw": WTW_JOINT_NAMES, "joint_names_sdk": list(SDK_JOINT_NAMES), } 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}") # ─── Ramp to default ─── def ramp_to_default(client, args, state): print("[INFO] Ramping to default pose (~2s)...") current_sdk = np.array([state.motorState[i].q for i in range(12)], dtype=np.float32) error = current_sdk - DEFAULT_ANGLES_SDK if np.max(np.abs(error)) < 0.05: print("[INFO] Already near default pose.") return state ramp_steps = 200 step_err = error / ramp_steps for i in range(ramp_steps): if EXIT: return state new_state = client.recv_latest() if new_state is not None: state = new_state targets = DEFAULT_ANGLES_SDK + (error - step_err * min(i + 1, ramp_steps)) cmd = LowCmd() for j in range(12): cmd.set_motor(j, MotorCmd( mode=MotorMode.Servo, q=float(targets[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) time.sleep(0.01) if i % 50 == 0: actual = np.array([state.motorState[j].q for j in range(12)]) print(f" ramp {i}/{ramp_steps} target_err={np.max(np.abs(targets-DEFAULT_ANGLES_SDK)):.3f} actual_err={np.max(np.abs(actual-DEFAULT_ANGLES_SDK)):.3f}") for _ in range(50): if EXIT: return state new_state = client.recv_latest() if new_state is not None: state = new_state send_hold_cmd(client, state, args) time.sleep(0.01) print("[INFO] Default pose reached.") return state # ─── Main ─── def main(): parser = argparse.ArgumentParser(description="WTW RMA deployment on Go1 PRO via go1_pro_sdk") parser.add_argument("--body", default=str(HERE / "body_latest.jit")) parser.add_argument("--adapt", default=str(HERE / "adaptation_module_latest.jit")) parser.add_argument("--kill-sport", action="store_true") parser.add_argument("--pi-host", default="192.168.123.161") parser.add_argument("--pi-user", default="pi") parser.add_argument("--kp", type=float, default=20.0) parser.add_argument("--kd", type=float, default=0.5) parser.add_argument("--kp-cal", type=float, default=15.0) parser.add_argument("--kd-cal", type=float, default=0.5) 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=1.0) parser.add_argument("--rc-vy-scale", type=float, default=1.0) parser.add_argument("--rc-wz-scale", type=float, default=1.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("--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("--log-dir", default="") parser.add_argument("--log-flush-every", type=int, default=50) args = parser.parse_args() print(""" ╔══════════════════════════════════════════════════════════════╗ ║ WTW RMA Deployment (go1_pro_sdk direct MCU) ║ ║ 1. Robot SUSPENDED ║ ║ 2. Use --kill-sport to auto-kill Pi processes ║ ║ 3. R2=go/stop, L2=estop, Left-stick=move, Right-stick=turn ║ ╚══════════════════════════════════════════════════════════════╝ """) input("Press Enter when ready...") if args.kill_sport: kill_sport_processes(args.pi_host, args.pi_user) print("[INFO] Loading WTW models...") policy = WTWPolicy(args.body, args.adapt) print("[INFO] Connecting to MCU...") client = MCUClient() logger = JsonlLogger(args.log_dir, args) try: print("[INFO] Waking MCU...") client.wake_mcu(n_frames=50, dt=0.01) state = client.recv_state(timeout=2.0) if state is None: print("[ERROR] No state received.") return 1 print(f"[INFO] Connected. Battery={state.bms.SOC}%") print(f"[INFO] RPY: {np.round(np.degrees(state.imu.rpy), 1)} deg") sm_state = State.IDLE edge = RCEdgeDetector() edge.update(state) clock = ClockState() base_cmd = build_commands_default() actions = np.zeros(NUM_ACTIONS, dtype=np.float32) last_actions = np.zeros(NUM_ACTIONS, dtype=np.float32) step = 0 dt = 1.0 / args.rate_hz next_t = time.perf_counter() print(f"[INFO] Rate: {args.rate_hz}Hz. Ctrl+C to exit.") print(f"[INFO] State machine: IDLE → (R2) → CALIBRATE → HOLD → (R2) → RL") 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, falling = edge.update(state) r2_rose = "R2" in rising l2_rose = "L2" in rising # Emergency stop if l2_rose and sm_state != State.IDLE: print(f"\n[L2 EMERGENCY] {sm_state.value} → IDLE") sm_state = State.IDLE actions[:] = 0.0 last_actions[:] = 0.0 policy.reset() clock.reset() client.send(LowCmd()) # State machine if sm_state == State.IDLE: if step % 10 == 0: client.send(LowCmd()) if r2_rose: print("\n[R2] IDLE → CALIBRATE") sm_state = State.CALIBRATE state = ramp_to_default(client, args, state) if EXIT: break 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 → RL") sm_state = State.RL actions[:] = 0.0 last_actions[:] = 0.0 policy.reset() clock.reset() elif sm_state == State.RL: if r2_rose: print("\n[R2] RL → HOLD") sm_state = State.HOLD actions[:] = 0.0 last_actions[:] = 0.0 state = ramp_to_default(client, args, state) if EXIT: break continue # Update commands from RC base_cmd = get_rc_commands(state, base_cmd, args) # Clock inputs clock_inputs = clock.step(base_cmd, dt) # Observation obs = compute_obs_wtw(state.imu, state.motorState, base_cmd, actions, last_actions, clock_inputs) # Inference action_raw = policy(obs) actions = np.clip(action_raw, -CLIP_ACTIONS, CLIP_ACTIONS).astype(np.float32) last_actions = actions.copy() if step >= args.warmup_steps: send_rl_cmd(client, state, action_raw, args) logger.log( step, mode="RL", commands=base_cmd, obs_wtw=obs, action_raw=action_raw, action_safe=actions, dof_pos_sdk=np.array([state.motorState[i].q for i in range(12)]), dof_vel_sdk=np.array([state.motorState[i].dq for i in range(12)]), clock_inputs=clock_inputs, imu_rpy_deg=np.degrees(state.imu.rpy), rc_buttons=state.remote.pressed, ) else: logger.log( step, mode=sm_state.value, dof_pos_sdk=np.array([state.motorState[i].q for i in range(12)]), imu_rpy_deg=np.degrees(state.imu.rpy), rc_buttons=state.remote.pressed, ) # Status print if step % args.print_every == 0: dof_pos = np.array([state.motorState[i].q for i in range(12)]) print(f"\n[STEP {step}] state={sm_state.value} bat={state.bms.SOC}% " f"rpy={np.round(np.degrees(state.imu.rpy), 1)}") print(f" RC: lx={state.remote.lx:+.2f} ly={state.remote.ly:+.2f} " f"btns={state.remote.pressed}") print(f" joint: {np.round(dof_pos, 2)}") if sm_state == State.RL: print(f" action max: {np.max(np.abs(actions)):.2f}") 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() finally: if logger is not None: logger.close() print("[INFO] Safe stopping...") client.safe_stop(n_frames=50, dt=0.002) client.close() print("[INFO] Done.") if __name__ == "__main__": main()