wtw,部署

This commit is contained in:
cyy_mac
2026-07-24 13:19:29 +08:00
parent 04eaa8c916
commit 554d4649c5
22 changed files with 1960 additions and 9 deletions

View File

@@ -0,0 +1,133 @@
#!/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()