Files
Motrixlab/scripts/view_orig.py

90 lines
3.8 KiB
Python

#!/usr/bin/env python3
"""MuJoCo viewer for original Go1 (45-dim, PD 80/1.0, action_scale=0.05)"""
import numpy as np, mujoco, onnxruntime as ort, os, time, threading, queue
from mujoco import viewer
from pynput import keyboard
PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
ONNX = os.path.join(PROJECT, "exports_go1_orig", "policy.onnx")
XML_DIR = os.path.join(PROJECT, "motrix_envs/src/motrix_envs/locomotion/go1/xmls")
# Original params
NUM_OBS = 57
KP, KD = 80.0, 0.5 # KD=0.5 + joint_damping(0.5) = 1.0 = training kd
ACTION_SCALE = 0.05
CLIP = 23.7
DEFAULT = 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)
class KB:
def __init__(self):
self._q=queue.Queue(); self.running=True; self.held=set()
def _n(self,k):
try:
if hasattr(k,'char') and k.char: return k.char.lower()
except: pass
return str(k).lower()
def _w(self):
while self.running:
try:
et,k=self._q.get(timeout=0.05); n=self._n(k)
if et=='press': self.held.add(n)
elif et=='release': self.held.discard(n)
except queue.Empty: pass
def init(self):
self._l=keyboard.Listener(on_press=lambda k:self._q.put(('press',k)), on_release=lambda k:self._q.put(('release',k)))
self._l.start(); self._t=threading.Thread(target=self._w,daemon=True); self._t.start()
def held_keys(self): return self.held.copy()
def stop(self): self.running=False; self._l.stop()
def get_sensor(m,d,name):
sid=mujoco.mj_name2id(m,mujoco.mjtObj.mjOBJ_SENSOR,name)
if sid<0: return None
adr=m.sensor_adr[sid]; return d.sensordata[adr:adr+m.sensor_dim[sid]].copy()
def compute_obs(model,data,cmd,last_a):
obs=np.zeros(NUM_OBS,dtype=np.float32)
g=get_sensor(model,data,"gyro")
obs[0:3]=(g if g is not None else data.qvel[3:6])*0.25
R=data.xmat[1].reshape(3,3)
obs[3:6]=(R.T@np.array([0.,0.,-1.])).astype(np.float32)
obs[6:18]=(data.qpos[7:19]-DEFAULT)*1.0
obs[18:30]=data.qvel[6:18]*0.05
obs[30:42]=last_a
obs[42:45]=cmd; obs[45:57]=0.0 # contact_force*np.array([2.,2.,0.25],dtype=np.float32)
return np.clip(obs,-100.,100.)
def main():
os.chdir(XML_DIR)
model=mujoco.MjModel.from_xml_string(open("scene_motor_actuator.xml").read())
data=mujoco.MjData(model)
data.qpos[0:3]=[0,0,0.42]; data.qpos[3:7]=[1,0,0,0]; data.qpos[7:19]=DEFAULT
mujoco.mj_forward(model,data)
session=ort.InferenceSession(ONNX,providers=['CPUExecutionProvider'])
print(f"[ORIG] kp={KP} kd={KD+0.5} action_scale={ACTION_SCALE} | W/S前后 Q/E左右 A/D旋转")
kb=KB(); kb.init()
view=viewer.launch_passive(model,data)
step,vx,vy,wz=0,0.,0.,0.
last_a=np.zeros(12,dtype=np.float32); action=np.zeros(12,dtype=np.float32)
dec=2
while view.is_running():
keys=kb.held_keys()
if 'escape' in keys: break
if 'r' in keys:
data.qpos[0:3]=[0,0,0.42]; data.qpos[3:7]=[1,0,0,0]; data.qpos[7:19]=DEFAULT
data.qvel[:]=0; last_a[:]=0; mujoco.mj_forward(model,data)
if ' ' in keys: vx=vy=wz=0.
vx=1.0 if 'w' in keys else (-1.0 if 's' in keys else 0.)
vy=1.0 if 'q' in keys else (-1.0 if 'e' in keys else 0.)
wz=1.0 if 'a' in keys else (-1.0 if 'd' in keys else 0.)
if step%dec==0:
obs=compute_obs(model,data,np.array([vx,vy,wz],dtype=np.float32),last_a)
action=session.run(None,{'observations':obs.reshape(1,-1).astype(np.float32)})[0][0]
action=np.clip(action,-CLIP,CLIP); last_a=action.copy()
target=DEFAULT+action*ACTION_SCALE
t=KP*(target-data.qpos[7:19])-KD*data.qvel[6:18]
data.ctrl[:]=np.clip(t,-CLIP,CLIP)
mujoco.mj_step(model,data); view.sync(); step+=1; time.sleep(0.001)
kb.stop(); view.close()
if __name__=="__main__": main()