Files

119 lines
10 KiB
Python

"""Prepare an L20 linear-mimic bottle task for the actual SPIDER contact-guided runner."""
import os
os.environ.setdefault('MUJOCO_GL','osmesa')
from pathlib import Path
import json,copy,xml.etree.ElementTree as E
import numpy as np
import mujoco
from scipy.spatial.transform import Rotation
from scipy.optimize import least_squares
import l20_model_source as source
ROOT=Path(__file__).resolve().parents[1];OUT=ROOT/'output/spider_l20_contact'
TASK=OUT/'datasets/processed/current/l20/right/bottle_hold';TRIAL=TASK/'0';TRIAL.mkdir(parents=True,exist_ok=True)
source.REPO_ROOT=ROOT/'third_party/l20_assets'
p=source.build_model('right',OUT);tree=E.parse(p);root=tree.getroot();world=root.find('worldbody');hand=world.find("body[@name='hand_base_link']")
for mesh in root.findall('asset/mesh'):
mesh.set('file',str((OUT/mesh.get('file')).resolve()));mesh.set('maxhullvert','64')
root.find('option').attrib.update(timestep='.0025',gravity='0 0 -9.81',integrator='implicitfast',cone='elliptic')
world.find("geom[@name='floor']").attrib.update(pos='0 0 0',contype='4',conaffinity='2')
for g in hand.iter('geom'):g.attrib.update(contype='1',conaffinity='2',friction='.6 .005 .001',condim='3',solref='.01 1')
root.find('default/joint').attrib.update(damping='.05',armature='.002')
# Scalar wrist coordinates are SPIDER's dexterous-hand control convention.
axes=['1 0 0','0 1 0','0 0 1'];wristnames=[]
for i in range(6):
name='right_hand_'+('pos_'+'xyz'[i] if i<3 else 'rot_'+'xyz'[i-3]);wristnames.append(name)
hand.insert(i,E.Element('joint',name=name,type='slide' if i<3 else 'hinge',axis=axes[i%3],limited='false',damping='0',armature='.02'))
urdf=E.parse(source.source_urdf()).getroot();joints=urdf.findall('joint');names=[j.get('name') for j in joints];active=[j.get('name') for j in joints if j.find('mimic') is None]
limits={j.get('name'):np.array([float(j.find('limit').get('lower')),float(j.find('limit').get('upper'))]) for j in joints};mimic={}
for j in joints:
mi=j.find('mimic')
if mi is not None:
name=j.get('name');a=mi.get('joint');scale=float(mi.get('multiplier','1'));offset=float(mi.get('offset','0'));mimic[name]=(a,scale,offset)
lo,hi=np.sort((limits[name]-offset)/scale);limits[a]=np.array([max(limits[a][0],lo),min(limits[a][1],hi)])
for eq in root.findall('equality/joint'):eq.attrib.update(solref='.01 1')
act=E.SubElement(root,'actuator')
for i,name in enumerate(wristnames):E.SubElement(act,'position',name=name+'_act',joint=name,kp='2000' if i<3 else '40',kv='60' if i<3 else '2',forcerange='-200 200' if i<3 else '-20 20')
for name in active:E.SubElement(act,'position',name=name+'_act',joint=name,kp='8',kv='.15',forcerange='-1 1',ctrlrange=source._fmt(limits[name]))
bottlepath=ROOT/'third_party/bottle_model_parametric/bottle.xml';bottle=E.parse(bottlepath).getroot()
for item in bottle.find('asset'):
item=copy.deepcopy(item)
if item.get('file'):item.set('file',str((bottlepath.parent/item.get('file')).resolve()))
if item.tag=='mesh':item.set('maxhullvert','64')
root.find('asset').append(item)
obj=E.SubElement(world,'body',name='right_object')
objectnames=[]
for i in range(6):
name='right_object_'+('pos_'+'xyz'[i] if i<3 else 'rot_'+'xyz'[i-3]);objectnames.append(name)
E.SubElement(obj,'joint',name=name,type='slide' if i<3 else 'hinge',axis=axes[i%3],limited='false',damping='0',armature='0')
E.SubElement(act,'position',name=name,joint=name,kp='0',kv='0')
for j,item in enumerate(bottle.find('worldbody/body')):
if item.tag=='freejoint':continue
item=copy.deepcopy(item)
if item.tag=='geom':
item.set('name',f'bottle_{j}')
if item.get('contype','1')!='0':item.attrib.update(contype='2',conaffinity='5',friction='.6 .005 .001',condim='3',solref='.01 1')
obj.append(item)
E.SubElement(obj,'site',name='right_object_track',pos='0 0 .07')
scene=TASK/'scene_act.xml';tree.write(scene)
model=mujoco.MjModel.from_xml_path(str(scene));data=mujoco.MjData(model)
addr={n:int(model.jnt_qposadr[model.joint(n).id]) for n in names+wristnames+objectnames}
actnames=[wristnames[i] for i in range(6)]+active+objectnames
source_motion=np.load(ROOT/'output/l20_2047635068/jitter_audit/temporal/motion.npz');start=600;stop=660
sourcenames=list(source_motion['joint_names']);wp=source_motion['wrist_pos'][start:stop];wr=Rotation.from_quat(source_motion['wrist_quat_wxyz'][start:stop][:,[1,2,3,0]])
# Estimated bottle registration, not an observed video object track.
relpos=np.array([.12,-.065,.15]);relrot=Rotation.from_euler('x',-np.pi/2)
op=wp+wr.apply(np.tile(relpos,(len(wp),1)));orr=wr*relrot
align=orr[0].inv();shift=np.array([0,0,.3])-align.apply(op[0]);wp=align.apply(wp)+shift;wr=align*wr;op=align.apply(op)+shift;orr=align*orr
baseq=np.stack([source_motion['qpos'][start:stop,sourcenames.index(n)] for n in active],axis=1)
low=np.array([limits[n][0] for n in active]);high=np.array([limits[n][1] for n in active]);baseq=np.clip(baseq,low,high)
handgeoms=np.flatnonzero(model.geom_contype==1);objectgeoms=np.flatnonzero(model.geom_contype==2)
tips=[next(g for g in handgeoms if model.geom_bodyid[g]==model.body(f+'_distal').id) for f in source.FINGERS]
def pose(a,pos,rot,objpos,objrot):
q=dict(zip(active,a));q.update({n:sc*q[parent]+off for n,(parent,sc,off) in mimic.items()})
for n,v in q.items():data.qpos[addr[n]]=v
data.qpos[[addr[n] for n in wristnames]]=np.r_[pos,rot.as_euler('XYZ')]
data.qpos[[addr[n] for n in objectnames]]=np.r_[objpos,objrot.as_euler('XYZ')]
mujoco.mj_kinematics(model,data)
def distances():return np.array([min(mujoco.mj_geomDistance(model,data,int(g),int(o),.5,None) for o in objectgeoms) for g in handgeoms])
ti=[list(handgeoms).index(g) for g in tips]
def residual(x):
pose(x[:16],wp[0]+x[16:19],wr[0]*Rotation.from_rotvec(x[19:]),op[0],orr[0]);d=distances()
return np.r_[(d[ti]+.0004)/.004,np.maximum(-d-.0005,0)/.001,.08*(x[:16]-baseq[0]),.2*x[16:19]/.03,.15*x[19:]/.3]
x0=np.r_[baseq[0],np.zeros(6)];residual(x0);before=distances()
fits=[]
for shift_x,shift_z in [(0,0),(.025,0),(-.025,0),(0,.025),(0,-.025),(.025,.025),(-.025,-.025)]:
seed=x0.copy();seed[16]=shift_x;seed[18]=shift_z
trial=least_squares(residual,seed,bounds=(np.r_[low,[-.06]*3,[-.5]*3],np.r_[high,[.06]*3,[.5]*3]),max_nfev=250,diff_step=1e-4)
residual(trial.x);ds=distances();score=float(np.abs(ds[ti]).mean()+3*max(0,-ds.min()-.001))
fits.append((score,trial));print('initialization seed',shift_x,shift_z,'gap mm',np.abs(ds[ti]).mean()*1000,'penetration mm',max(0,-ds.min())*1000,flush=True)
fit=min(fits,key=lambda x:x[0])[1]
residual(fit.x);after=distances();contactlocal=[]
# Sites mark the actual closest point on each distal link at initialization.
for finger,g in zip(source.FINGERS,tips):
candidates=[]
for o in objectgeoms:
segment=np.zeros(6);dist=mujoco.mj_geomDistance(model,data,int(g),int(o),.5,segment);candidates.append((dist,segment.copy()))
_,segment=min(candidates,key=lambda x:x[0]);bid=model.geom_bodyid[g];local=data.xmat[bid].reshape(3,3).T@(segment[:3]-data.xpos[bid])
link=hand.find(f".//body[@name='{finger}_distal']");E.SubElement(link,'site',name=f'right_hand_{finger}_track',pos=source._fmt(local))
contactlocal.append(orr[0].inv().apply(segment[3:]-op[0]))
tree.write(scene);model=mujoco.MjModel.from_xml_path(str(scene));data=mujoco.MjData(model)
contacts=[model.site(f'right_hand_{f}_track').id for f in source.FINGERS]
# Preserve source movement, applying one smooth initial registration correction.
source_times=np.arange(stop-start)/30;times=np.arange(801)*.0025;rows=[];controls=[];contactpos=[]
from scipy.spatial.transform import Slerp
wpi=np.stack([np.interp(times,source_times,wp[:,k]) for k in range(3)],axis=1)
opi=np.stack([np.interp(times,source_times,op[:,k]) for k in range(3)],axis=1)
wri=Slerp(source_times,wr)(np.minimum(times,source_times[-1]));ori=Slerp(source_times,orr)(np.minimum(times,source_times[-1]))
ai=np.stack([np.interp(times,source_times,baseq[:,k]) for k in range(16)],axis=1);ai=np.clip(ai+fit.x[:16]-baseq[0],low,high)
for t in range(len(times)):
pose(ai[t],wpi[t]+fit.x[16:19],wri[t]*Rotation.from_rotvec(fit.x[19:]),opi[t],ori[t])
rows.append(data.qpos.copy());controls.append([data.qpos[addr[n]] for n in actnames]);contactpos.append(ori[t].apply(contactlocal)+opi[t])
q=np.array(rows);ctrl=np.array(controls);vel=np.gradient(q,.0025,axis=0);vel[0]=0
np.savez_compressed(TRIAL/'trajectory_kinematic_act.npz',qpos=q,qvel=vel,ctrl=ctrl,contact=np.ones((len(q),5)),contact_pos=contactpos,time=times)
(TASK/'task_info.json').write_text(json.dumps(dict(ref_dt=.0025,contact_site_ids=contacts),indent=2))
# The final real rollout has zero object actuator gains; virtual assistance only during search.
config=dict(dataset_dir=str(OUT/'datasets'),dataset_name='current',robot_type='l20',embodiment_type='right',task='bottle_hold',data_id=0,contact_guidance=True,contact_rew_scale=1.,device='cuda:0',sim_dt=.0025,ref_dt=.0025,ctrl_dt=.1,knot_dt=.1,horizon=.6,trace_dt=.01,render_dt=.02,num_samples=512,max_num_iterations=6,num_dr=1,num_dyn=1,pair_margin_range=[0.,0.],xy_offset_range=[0.,0.],guidance_decay_ratio=.4,init_pos_actuator_gain=10.,init_pos_actuator_bias=1.,init_rot_actuator_gain=.1,init_rot_actuator_bias=.02,max_sim_steps=800,show_viewer=False,viewer='',wait_on_finish=False,save_video=False,save_info=True,save_metrics=False,save_config=True,sanity_check_seconds=0.,use_torch_compile=False,improvement_threshold=0.,pos_noise_scale=.008,rot_noise_scale=.05,joint_noise_scale=.2,exploit_ratio=.02,base_pos_rew_scale=.3,base_rot_rew_scale=.1,joint_rew_scale=.003,pos_rew_scale=5.,rot_rew_scale=.3,vel_rew_scale=.0001,temperature=.1,nconmax_per_env=128,njmax_per_env=512)
(OUT/'config.json').write_text(json.dumps(config,indent=2))
report=dict(source_frames=[start,stop],duration_s=2.,source='Current video trajectory, offline smoothed',bottle_alignment='Estimated constant bottle-to-palm relation, not tracked from video',coupling='Original URDF linear mimic, no nonlinear JSON adaptor',active_finger_joints=16,passive_finger_joints=5,controlled_wrist_dofs=6,object_virtual_actuators=6,object_real_rollout_gains_zero=True,initial_fit_converged=bool(fit.success),initial_mean_surface_gap_before_mm=float(np.abs(before[ti]).mean()*1000),initial_mean_surface_gap_after_mm=float(np.abs(after[ti]).mean()*1000),initial_max_penetration_mm=float(max(0,-after.min())*1000),nq=model.nq,nv=model.nv,nu=model.nu,contact_sites=contacts)
(OUT/'preparation.json').write_text(json.dumps(report,indent=2));print(json.dumps(report,indent=2))