Files
go1_pro_deploy/sim2sim_test_deploy.py
2026-06-21 21:58:39 +08:00

415 lines
14 KiB
Python

#!/usr/bin/env python3
"""
sim2sim test for deploy_onnx_pro_sdk.py using sim2sim_mujoco_example's Go1 XML.
Mirrors the deploy safety layer (joint limits, torque factor, position deviation protect)
and logs all data to JSONL when --log-dir is set.
Usage:
conda activate free_dog_sdk
mjpython sim2sim_test_deploy.py --onnx policy.onnx
mjpython sim2sim_test_deploy.py --onnx policy.onnx --log-dir logs
Controls: W/S=forward/back, Q/E=strafe, A/D=rotate, Space=stop, R=reset, Esc=quit
"""
import argparse
import json
import os
import signal
import time
from datetime import datetime
from pathlib import Path
import mujoco
import numpy as np
import onnxruntime as ort
from mujoco import viewer
HERE = os.path.dirname(os.path.abspath(__file__))
SIM2SIM_XML = os.path.join(HERE, "sim2sim_mujoco_example", "data", "go1", "xml", "go1.xml")
# ─── Policy constants (matches go1_sim2sim.py + deploy_onnx_pro_sdk.py) ───
NUM_OBS = 45
NUM_ACTIONS = 12
ACTION_SCALE = 0.05
CLIP_ACTIONS = 23.7
CLIP_OBS = 100.0
KP_DEFAULT = 80.0
KD_DEFAULT = 0.5 # + joint_damping(0.5) = 1.0 total
DEFAULT_ANGLES = np.array([
-0.0, 0.9, -1.8, # FR
0.0, 0.9, -1.8, # FL
-0.0, 0.9, -1.8, # RR
0.0, 0.9, -1.8, # RL
], dtype=np.float32)
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",
]
# ─── Safety limits (mirrors deploy go1_pro_sdk safety layer) ───
JOINT_TYPE = [("hip" if i % 3 == 0 else "thigh" if i % 3 == 1 else "knee") for i in range(12)]
JOINT_LIMITS = {
"hip": (-0.78, 0.78),
"thigh": (-0.60, 3.50),
"knee": (-2.70, -0.95),
}
TAU_MAX = {
"hip": 23.7,
"thigh": 23.7,
"knee": 35.55,
}
EXIT = False
def _sig_handler(signum, frame):
global EXIT
EXIT = True
signal.signal(signal.SIGINT, _sig_handler)
signal.signal(signal.SIGTERM, _sig_handler)
# ─── Safety functions (mirror deploy apply_safety) ───
def clip_targets_to_limits(targets):
"""PositionLimit: clamp target joint angles to JOINT_LIMITS."""
n_clamped = 0
safe = targets.copy()
for i in range(12):
lo, hi = JOINT_LIMITS[JOINT_TYPE[i]]
if safe[i] < lo:
safe[i] = lo; n_clamped += 1
elif safe[i] > hi:
safe[i] = hi; n_clamped += 1
return safe, n_clamped
def clip_torques(torques, power_factor):
"""PowerProtect: clamp torque to TAU_MAX * power_factor / 10."""
tau_lim = np.array([TAU_MAX[JOINT_TYPE[i]] * power_factor / 10.0 for i in range(12)],
dtype=np.float64)
return np.clip(torques, -tau_lim, tau_lim)
def position_protect_mask(targets, current_pos, limit_rad):
"""PositionProtect: return bool mask, True where deviation <= limit."""
return np.abs(targets - current_pos) <= limit_rad
# ─── Quaternion math (same as deploy_onnx_pro_sdk.py) ───
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)
# ─── Observations (same as deploy_onnx_pro_sdk.py) ───
def compute_obs(model, data, commands, last_actions):
obs = np.zeros(NUM_OBS, dtype=np.float32)
sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, "Body_Gyro")
adr = model.sensor_adr[sid]
obs[0:3] = data.sensordata[adr:adr + 3] * 0.25
sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, "Body_Quat")
adr = model.sensor_adr[sid]
quat = data.sensordata[adr:adr + 4]
R = quat_to_rot_matrix(quat)
obs[3:6] = (R.T @ np.array([0., 0., -1.], dtype=np.float64)).astype(np.float32)
dof_pos = np.zeros(12, dtype=np.float32)
dof_vel = np.zeros(12, dtype=np.float32)
for i, name in enumerate(JOINT_NAMES):
sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, f"{name}_pos")
dof_pos[i] = data.sensordata[model.sensor_adr[sid]]
sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, f"{name}_vel")
dof_vel[i] = data.sensordata[model.sensor_adr[sid]]
obs[6:18] = (dof_pos - DEFAULT_ANGLES) * 1.0
obs[18:30] = dof_vel * 0.05
obs[30:42] = last_actions
obs[42:45] = np.array(commands, dtype=np.float32) * np.array([2.0, 2.0, 0.25], dtype=np.float32)
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
# ─── 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"sim2sim_run_{ts}"
self.run_dir.mkdir(parents=True, exist_ok=True)
meta = {
"created_at": ts,
"args": vars(args),
"num_obs": NUM_OBS,
"num_actions": NUM_ACTIONS,
"action_scale": ACTION_SCALE,
"default_angles": DEFAULT_ANGLES.tolist(),
"joint_names": JOINT_NAMES,
"joint_limits": {k: list(v) for k, v in JOINT_LIMITS.items()},
"tau_max": TAU_MAX,
}
(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")
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}")
# ─── Keyboard input ───
class Keyboard:
def __init__(self):
self.held = set()
def _on_press(self, key):
try: self.held.add(key.char.lower())
except AttributeError: self.held.add(str(key))
def _on_release(self, key):
try: self.held.discard(key.char.lower())
except AttributeError: self.held.discard(str(key))
def init(self):
from pynput import keyboard
self._listener = keyboard.Listener(
on_press=self._on_press, on_release=self._on_release)
self._listener.start()
def keys(self):
return self.held.copy()
def stop(self):
if self._listener:
self._listener.stop()
# ─── Main ───
def main():
parser = argparse.ArgumentParser(description="Sim2sim test for deploy_onnx_pro_sdk.py")
parser.add_argument("--onnx", default=os.path.join(HERE, "policy.onnx"))
# PD gains
parser.add_argument("--kp", type=float, default=KP_DEFAULT)
parser.add_argument("--kd", type=float, default=KD_DEFAULT)
# Safety (mirrors deploy args)
parser.add_argument("--power-factor", type=int, default=7,
help="Torque limit factor 1-10, applied as TAU_MAX * factor/10")
parser.add_argument("--position-protect-limit", type=float, default=0.5,
help="Max |target - actual| before zeroing torque (negative=disable)")
parser.add_argument("--no-joint-limit", action="store_true",
help="Disable joint limit clipping on targets")
# Logging
parser.add_argument("--log-dir", default="", help="Enable JSONL logging to this directory")
parser.add_argument("--print-every", type=int, default=200)
args = parser.parse_args()
if not os.path.exists(SIM2SIM_XML):
print(f"[ERROR] XML not found: {SIM2SIM_XML}"); return 1
if not os.path.exists(args.onnx):
print(f"[ERROR] ONNX not found: {args.onnx}"); return 1
print(f"[INFO] XML: {SIM2SIM_XML}")
print(f"[INFO] ONNX: {args.onnx}")
model = mujoco.MjModel.from_xml_path(SIM2SIM_XML)
data = mujoco.MjData(model)
model.dof_damping[6:] = 0.5
total_kd = args.kd + model.dof_damping[6]
print(f"[INFO] Bodies={model.nbody}, DoF={model.nq}, Actuators={model.nu}")
print(f"[INFO] Timestep={model.opt.timestep}")
print(f"[INFO] KP={args.kp}, KD(active)={args.kd}, KD(passive)={model.dof_damping[6]}, "
f"total_KD={total_kd}")
print(f"[INFO] Safety: power_factor={args.power_factor}, "
f"position_protect={args.position_protect_limit}, "
f"joint_limit={not args.no_joint_limit}")
print(f"[INFO] Torque limits (factor={args.power_factor}): "
+ ", ".join(f"{jt}={TAU_MAX[jt]*args.power_factor/10:.1f}" for jt in ["hip","thigh","knee"]))
# Init pose
data.qpos[0:3] = [0.0, 0.0, 0.42]
data.qpos[3:7] = [1.0, 0.0, 0.0, 0.0]
data.qpos[7:19] = DEFAULT_ANGLES
data.qvel[:] = 0.0
mujoco.mj_forward(model, data)
# Load ONNX
session = ort.InferenceSession(args.onnx, providers=["CPUExecutionProvider"])
input_name = session.get_inputs()[0].name
print(f"[INFO] ONNX input={input_name}, shape={session.get_inputs()[0].shape}")
print(f"[INFO] Controls: W/S=前后 Q/E=左右 A/D=旋转 Space=停 R=重置 Esc=退出")
logger = JsonlLogger(args.log_dir, args)
kb = Keyboard(); kb.init()
view = viewer.launch_passive(model, data)
step = 0
ctrl_dt = 0.01
steps_per_inference = int(ctrl_dt / model.opt.timestep)
last_actions = np.zeros(NUM_ACTIONS, dtype=np.float32)
action = np.zeros(NUM_ACTIONS, dtype=np.float32)
safety_stats = {"joint_limit_clamps": 0, "position_protect_hits": 0}
t0 = time.perf_counter()
while view.is_running() and not EXIT:
t_loop = time.perf_counter()
keys = kb.keys()
if 'key.esc' in keys or '\x1b' in keys:
break
if 'r' in keys:
data.qpos[0:3] = [0.0, 0.0, 0.42]
data.qpos[3:7] = [1.0, 0.0, 0.0, 0.0]
data.qpos[7:19] = DEFAULT_ANGLES
data.qvel[:] = 0.0
last_actions[:] = 0.0
action[:] = 0.0
mujoco.mj_forward(model, data)
vx = 1.0 if 'w' in keys else (-1.0 if 's' in keys else 0.0)
vy = 1.0 if 'q' in keys else (-1.0 if 'e' in keys else 0.0)
wz = 0.5 if 'a' in keys else (-0.5 if 'd' in keys else 0.0)
if ' ' in keys:
vx = vy = wz = 0.0
commands = np.array([vx, vy, wz], dtype=np.float32)
# Inference at 100 Hz
if step % steps_per_inference == 0:
obs = compute_obs(model, data, commands, last_actions)
action = session.run(None, {input_name: obs.reshape(1, -1).astype(np.float32)})[0][0]
action = np.clip(action, -CLIP_ACTIONS, CLIP_ACTIONS)
last_actions = action.copy()
# Targets (pre-safety)
targets_raw = DEFAULT_ANGLES + action * ACTION_SCALE
# ── Safety layer (mirrors deploy apply_safety) ──
# 1. Joint limit clipping
if not args.no_joint_limit:
targets, n_clamped = clip_targets_to_limits(targets_raw)
safety_stats["joint_limit_clamps"] += n_clamped
else:
targets = targets_raw
current_pos = data.qpos[7:19]
current_vel = data.qvel[6:18]
# 2. Position deviation protection (zero torque where |target - actual| > limit)
pos_ok = np.ones(12, dtype=bool)
if args.position_protect_limit > 0:
pos_ok = position_protect_mask(targets, current_pos, args.position_protect_limit)
n_hit = np.sum(~pos_ok)
safety_stats["position_protect_hits"] += n_hit
# 3. PD control
torques = np.zeros(12, dtype=np.float64)
torques[pos_ok] = (args.kp * (targets[pos_ok] - current_pos[pos_ok])
- args.kd * current_vel[pos_ok])
# 4. Torque limiting (power_protect)
torques = clip_torques(torques, args.power_factor)
data.ctrl[:] = torques
mujoco.mj_step(model, data)
view.sync()
# Logging
logger.log(
step,
mode="rl",
loop_ms=(time.perf_counter() - t_loop) * 1000.0,
commands=commands,
obs=obs if step % steps_per_inference == 0 else np.zeros(0),
action_raw=action,
action_safe=action,
joint_targets_raw=targets_raw,
joint_targets_safe=targets,
dof_pos=current_pos,
dof_vel=current_vel,
torques=torques,
position_protect_hit_mask=(~pos_ok).astype(int),
)
step += 1
if step % args.print_every == 0:
z = data.qpos[2]
sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, "Body_Quat")
quat = data.sensordata[model.sensor_adr[sid]:model.sensor_adr[sid] + 4]
pos_err = np.max(np.abs(targets - current_pos))
n_pp = safety_stats["position_protect_hits"]
n_jl = safety_stats["joint_limit_clamps"]
print(f"\n[STEP {step}] z={z:.3f} cmd=[{vx:.1f},{vy:.1f},{wz:.1f}] "
f"max_err={pos_err:.3f}")
print(f" quat={np.round(quat, 3)} action_max={np.max(np.abs(action)):.2f}")
print(f" target: {np.round(targets[:4], 2)}")
print(f" actual: {np.round(current_pos[:4], 2)}")
print(f" torque: {np.round(torques[:4], 2)}")
if n_pp > 0 or n_jl > 0:
print(f" safety: pos_protect_hits={n_pp} joint_limit_clamps={n_jl}")
# Real-time sync
expected = (step + 1) * model.opt.timestep
elapsed = time.perf_counter() - t0
sleep = expected - elapsed
if sleep > 0:
time.sleep(sleep)
logger.close()
kb.stop()
view.close()
print("[INFO] Done.")
if __name__ == "__main__":
main()