Files
go1_pro_deploy/deploy_wtw/deploy_wtw_pro_sdk.py
2026-07-24 13:19:29 +08:00

651 lines
22 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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()