mp_yam_code / scripts /yam_kinematics.py
yqi19's picture
YAM bimanual task suite: env, solvers, tasks, converters
7399b6f verified
Raw
History Blame Contribute Delete
4.55 kB
"""Extract YAM right-arm kinematics (link chain: origin transforms, axes, sizes) -> JSON for a web FK viewer."""
import argparse, sys, os, json
from isaaclab.app import AppLauncher
p=argparse.ArgumentParser(); AppLauncher.add_app_launcher_args(p); a=p.parse_args(); a.headless=True; a.enable_cameras=False
app=AppLauncher(a).app
import numpy as np, torch, gymnasium as gym
REPO=os.path.dirname(os.path.dirname(os.path.abspath(__file__))); sys.path.insert(0,os.path.join(REPO,"source"))
import bimanual.tasks.manager_based.yam # noqa
from isaaclab_tasks.utils import parse_env_cfg
TASK="Template-YAM-Play-v0"; dev="cuda:0"
env=gym.make(TASK, cfg=parse_env_cfg(TASK, device=dev, num_envs=1)); u=env.unwrapped; env.reset()
R=u.scene["right_robot"]; bn=list(R.data.body_names); jn=list(R.data.joint_names)
origin=u.scene.env_origins[0].cpu().numpy()
def Rq(q):
w,x,y,z=q; return np.array([[1-2*(y*y+z*z),2*(x*y-z*w),2*(x*z+y*w)],[2*(x*y+z*w),1-2*(x*x+z*z),2*(y*z-x*w)],[2*(x*z-y*w),2*(y*z+x*w),1-2*(x*x+y*y)]])
def T(p,q):
M=np.eye(4); M[:3,:3]=Rq(q); M[:3,3]=p; return M
def set_joints(vals):
q=R.data.joint_pos.clone();
for i in range(6): q[0,jn.index(f"joint{i+1}")]=vals[i]
R.write_joint_state_to_sim(q, torch.zeros_like(q))
for _ in range(2): u.sim.step(); u.scene.update(1/60)
def bodyT(name):
i=bn.index(name); p=R.data.body_pos_w[0,i].cpu().numpy()-origin; q=R.data.body_quat_w[0,i].cpu().numpy(); return T(p,q)
def bbox(name):
try:
import omni.usd; from pxr import UsdGeom, Usd
stage=omni.usd.get_context().get_stage()
pp=R.root_physx_view.prim_paths[0] # articulation root; link prims are children
prim=stage.GetPrimAtPath(pp+"/"+name) if stage.GetPrimAtPath(pp+"/"+name).IsValid() else None
if prim is None:
# search by name under root
for pr in Usd.PrimRange(stage.GetPrimAtPath(pp)):
if pr.GetName()==name: prim=pr; break
bb=UsdGeom.BBoxCache(Usd.TimeCode.Default(),[UsdGeom.Tokens.default_,UsdGeom.Tokens.render])
rng=bb.ComputeLocalBound(prim).ComputeAlignedRange()
mn=np.array(rng.GetMin()); mx=np.array(rng.GetMax()); return [(mx-mn).tolist(),(0.5*(mn+mx)).tolist()]
except Exception as e:
return [[0.06,0.06,0.08],[0,0,0]]
def _q(M):
m=M[:3,:3]; t=m[0,0]+m[1,1]+m[2,2]
if t>0: s=np.sqrt(t+1)*2; w=.25*s; x=(m[2,1]-m[1,2])/s; y=(m[0,2]-m[2,0])/s; z=(m[1,0]-m[0,1])/s
elif m[0,0]>m[1,1] and m[0,0]>m[2,2]: s=np.sqrt(1+m[0,0]-m[1,1]-m[2,2])*2; w=(m[2,1]-m[1,2])/s; x=.25*s; y=(m[0,1]+m[1,0])/s; z=(m[0,2]+m[2,0])/s
elif m[1,1]>m[2,2]: s=np.sqrt(1+m[1,1]-m[0,0]-m[2,2])*2; w=(m[0,2]-m[2,0])/s; x=(m[0,1]+m[1,0])/s; y=.25*s; z=(m[1,2]+m[2,1])/s
else: s=np.sqrt(1+m[2,2]-m[0,0]-m[1,1])*2; w=(m[1,0]-m[0,1])/s; x=(m[0,2]+m[2,0])/s; y=(m[1,2]+m[2,1])/s; z=.25*s
q=np.array([w,x,y,z]); return q/(np.linalg.norm(q)+1e-9)
chain=["link_1","link_2","link_3","link_4","link_5","link_6"]
parent={"link_1":"arm","link_2":"link_1","link_3":"link_2","link_4":"link_3","link_5":"link_4","link_6":"link_5"}
set_joints([0,0,0,0,0,0])
T0={n:bodyT(n) for n in ["arm"]+chain}
links=[]
for i,n in enumerate(chain):
par=parent[n]
origin_T=np.linalg.inv(T0[par])@T0[n] # child origin relative to parent at zero
# axis: perturb joint i by +0.3, recompute local rotation of link n rel parent
v=[0]*6; v[i]=0.3; set_joints(v)
Tp=np.linalg.inv(bodyT(par))@bodyT(n)
set_joints([0]*6)
Rrel=np.linalg.inv(origin_T[:3,:3])@Tp[:3,:3] # = Rot(axis,0.3)
ang=np.arccos(np.clip((np.trace(Rrel)-1)/2,-1,1))
if ang>1e-4:
ax=np.array([Rrel[2,1]-Rrel[1,2],Rrel[0,2]-Rrel[2,0],Rrel[1,0]-Rrel[0,1]]); ax=ax/(np.linalg.norm(ax)+1e-9)
else: ax=np.array([0,0,1.0])
sz,ctr=bbox(n)
links.append({"name":n,"parent":par,
"origin_pos":origin_T[:3,3].tolist(),
"origin_quat":[float(x) for x in _q(origin_T)],
"axis":ax.tolist(),"size":sz,"center":ctr})
print(f"[k] {n} parent={par} axis={np.round(ax,2)} size={np.round(sz,3)}", flush=True)
base=T0["arm"]
out={"base_pos":base[:3,3].tolist(),"base_quat":[float(x) for x in _q(base)],
"links":links,
"limits":[[-2.618,3.054],[0.0,3.665],[0.0,3.665],[-1.571,1.571],[-1.571,1.571],[-2.094,2.094]],
"home":[0.0,1.5708,0.7854,0.7854,0.0,0.0]}
json.dump(out, open("/home/horde/project/xiaotong/outputs/yam_web/yam_kinematics.json","w"))
print("[k] saved kinematics json", flush=True)
env.close(); app.close(); print("KIN_OK", flush=True)