#!/usr/bin/env python3 """ Sim2sim test for 57-dim no-linevel policy (better3.onnx). obs: gyro(3) + gravity(3) + dof_pos(12) + dof_vel(12) + last_actions(12) + commands(3) + contacts(12) = 57 Contacts set to zero. Usage: conda activate free_dog_sdk mjpython sim2sim_57dim_test.py mjpython sim2sim_57dim_test.py --action-scale 0.07 --action-ema-alpha 0.3 Controls: W/S=前后 Q/E=左右 A/D=旋转 Space=停 R=重置 Esc=退出 """ import argparse, os, signal, time import mujoco, numpy as np, onnxruntime as ort from mujoco import viewer HERE = os.path.dirname(os.path.abspath(__file__)) XML = os.path.join(HERE, "..", "sim2sim_mujoco_example", "data", "go1", "xml", "go1.xml") ONNX = os.path.join(HERE, "better3.onnx") NUM_OBS, NUM_ACTIONS = 57, 12 DEFAULT_ANGLES = np.array([-0.0,0.9,-1.8, 0.0,0.9,-1.8, -0.0,0.9,-1.8, 0.0,0.9,-1.8], dtype=np.float32) EXIT = False def _sig(s, f): global EXIT; EXIT = True signal.signal(signal.SIGINT, _sig) class KB: def __init__(s): s.h = set(); s._l = None def _p(s,k): try: s.h.add(k.char.lower()) except: s.h.add(str(k)) def _r(s,k): try: s.h.discard(k.char.lower()) except: s.h.discard(str(k)) def init(s): from pynput import keyboard s._l = keyboard.Listener(on_press=s._p, on_release=s._r); s._l.start() def keys(s): return s.h.copy() def stop(s): if s._l: s._l.stop() def main(): p = argparse.ArgumentParser() p.add_argument("--action-scale", type=float, default=0.05) p.add_argument("--action-ema-alpha", type=float, default=0.0) p.add_argument("--kp", type=float, default=80.0) p.add_argument("--kd", type=float, default=1.0) args = p.parse_args() model = mujoco.MjModel.from_xml_path(XML) data = mujoco.MjData(model) print(f"[INFO] 57-dim no-linevel ONNX: {ONNX}") print(f"[INFO] action_scale={args.action_scale} ema_alpha={args.action_ema_alpha}") print(f"[INFO] KP={args.kp} KD={args.kd} passive_damping={model.dof_damping[6]}") sess = ort.InferenceSession(ONNX, providers=["CPUExecutionProvider"]) input_name = sess.get_inputs()[0].name kb = KB(); kb.init() data.qpos[:3] = [0,0,0.42]; data.qpos[3:7] = [1,0,0,0]; data.qpos[7:19] = DEFAULT_ANGLES mujoco.mj_forward(model, data) view = viewer.launch_passive(model, data) step = 0; si = int(0.01 / model.opt.timestep) # 100Hz last_a = np.zeros(12, dtype=np.float32); action = np.zeros(12, dtype=np.float32) t0 = time.perf_counter() while view.is_running() and not EXIT: keys = kb.keys() if 'key.esc' in keys: break if 'r' in keys: data.qpos[:3]=[0,0,0.42]; data.qpos[3:7]=[1,0,0,0] data.qpos[7:19]=DEFAULT_ANGLES; data.qvel[:]=0; last_a[:]=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=1.0 if 'a' in keys else (-1.0 if 'd' in keys else 0.0) if ' ' in keys: vx=vy=wz=0.0 if step % si == 0: obs = np.zeros(NUM_OBS, dtype=np.float32) # gyro sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, "Body_Gyro") obs[0:3] = data.sensordata[model.sensor_adr[sid]:model.sensor_adr[sid]+3] * 0.25 # gravity R = data.xmat[1].reshape(3, 3) obs[3:6] = (R.T @ [0., 0., -1.]).astype(np.float32) # dof_pos obs[6:18] = (data.qpos[7:19] - DEFAULT_ANGLES) * 1.0 # dof_vel obs[18:30] = data.qvel[6:18] * 0.05 # last_actions obs[30:42] = last_a # commands obs[42:45] = np.array([vx, vy, wz]) * np.array([2., 2., 0.25]) # contact forces = 0 obs[45:57] = 0.0 obs = np.clip(obs, -100., 100.) action = sess.run(None, {input_name: obs.reshape(1, -1).astype(np.float32)})[0][0] action = np.clip(action, -23.7, 23.7) if 0 < args.action_ema_alpha < 1: action = args.action_ema_alpha * action + (1 - args.action_ema_alpha) * last_a last_a = action.copy() targets = DEFAULT_ANGLES + action * args.action_scale data.ctrl[:] = np.clip(args.kp*(targets-data.qpos[7:19])-args.kd*data.qvel[6:18], -23.7, 23.7) mujoco.mj_step(model, data) view.sync() if step % 200 == 0: print(f"[{step}] z={data.qpos[2]:.3f} cmd=[{vx:.1f},{vy:.1f},{wz:.1f}] " f"pos=[{data.qpos[0]:.2f},{data.qpos[1]:.2f}]") step += 1 expected = (step+1)*model.opt.timestep sleep = expected - (time.perf_counter()-t0) if sleep > 0: time.sleep(sleep) kb.stop(); view.close() if __name__ == "__main__": main()