deploy 成功
This commit is contained in:
414
sim2sim_test_deploy.py
Normal file
414
sim2sim_test_deploy.py
Normal file
@@ -0,0 +1,414 @@
|
||||
#!/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()
|
||||
Reference in New Issue
Block a user