415 lines
14 KiB
Python
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()
|