DynaTTT / scripts /eval_widowx_simpler.py
DanTim05's picture
Upload Zeva source and project assets
68eed61 verified
Raw History Blame Contribute Delete
11.8 kB
"""Configurable-count, deterministic-seed SimplerEnv episodes via localhost."""
import argparse
import asyncio
import json
import math
import subprocess
from pathlib import Path
import cv2
import msgpack
import numpy as np
from scipy.spatial.transform import Rotation
from transforms3d.euler import mat2euler
import websockets
import simpler_env
from widowx_rollout import CausalHistory
from scripts.bridge_recovery_contract import legacy_normalizer, reached_delta_raw
from simpler_env.utils.env.observation_utils import get_image_from_maniskill2_obs_dict
_LEGACY_NORMALIZER = legacy_normalizer()
Q01 = np.asarray(_LEGACY_NORMALIZER['q01'], dtype=np.float32)
Q99 = np.asarray(_LEGACY_NORMALIZER['q99'], dtype=np.float32)
def pack(x):
if isinstance(x,np.ndarray): return {b'__ndarray__':True,b'data':x.tobytes(),b'dtype':x.dtype.str,b'shape':x.shape}
if isinstance(x,np.generic): return x.item()
raise TypeError(type(x))
def unpack(x):
if b'__ndarray__' in x: return np.frombuffer(x[b'data'],dtype=x[b'dtype']).reshape(x[b'shape'])
return x
def proprio(env, convention='bridge_tool'):
e=env.unwrapped
pose=e.agent.robot.pose.inv()*e.tcp.pose
# Bridge state uses the tool-frame convention from the dataset replay,
# not SimplerEnv's direct TCP Euler angles. The fixed basis transform is
# R_sim = R_bridge @ mat_transform; invert it before mat2euler.
sim_rotation=np.asarray(pose.to_transformation_matrix(),dtype=np.float64)[:3,:3]
mat_transform=np.asarray([[0.0,0.0,1.0],[0.0,1.0,0.0],[-1.0,0.0,0.0]],dtype=np.float64)
if convention == 'bridge_tool':
angles=np.asarray(mat2euler(sim_rotation @ mat_transform.T,axes='sxyz'),dtype=np.float32)
elif convention == 'legacy_tcp':
q=pose.q
angles=Rotation.from_quat([q[1],q[2],q[3],q[0]]).as_euler('xyz')
else:
raise ValueError(f'Unknown proprio convention: {convention}')
openness=1-float(e.agent.get_gripper_closedness())
return np.asarray([*pose.p,*angles,0,openness,openness],np.float32)
async def evaluate(args):
if args.proprio_convention != 'bridge_tool':
raise ValueError('Reached-delta CTE history requires the Bridge tool-frame proprio convention')
results=[]
async with websockets.connect(f'ws://127.0.0.1:{args.port}',compression=None,max_size=None,ping_interval=None) as ws:
metadata=msgpack.unpackb(await ws.recv(),object_hook=unpack)
if metadata.get('memory_mode') != args.memory_mode:
raise RuntimeError(f'Policy memory mode mismatch: {metadata} vs {args.memory_mode}')
for episode in range(args.seed_start, args.seed_start + args.episodes):
env=simpler_env.make('widowx_carrot_on_plate',
renderer_kwargs=dict(device=args.render_device, offscreen_only=True))
max_steps = args.max_steps
if max_steps is None and args.max_seconds is not None:
# SimplerEnv's WidowX control loop is 5 Hz. Override only the
# outer Gymnasium TimeLimit so the policy/evaluation protocol
# remains unchanged while allowing a longer episode horizon.
max_steps = max(1, int(math.ceil(args.max_seconds * 5.0)))
if max_steps is not None:
if not hasattr(env, '_max_episode_steps'):
raise RuntimeError('expected a Gymnasium TimeLimit wrapper')
env._max_episode_steps = max_steps
obs,_=env.reset(seed=episode)
image=lambda o: cv2.resize(get_image_from_maniskill2_obs_dict(env,o),(256,256))
first_image=image(obs); history=CausalHistory(first_image); steps=0; success=False; done=False
progress={}; trace=[]
writer = None
video_path = None
if args.video_dir:
video_dir=Path(args.video_dir); video_dir.mkdir(parents=True,exist_ok=True)
video_path=video_dir/f"{Path(args.checkpoint).stem}_episode_{episode:02d}.mp4"
writer=cv2.VideoWriter(str(video_path),cv2.VideoWriter_fourcc(*'mp4v'),5.0,(256,256))
if not writer.isOpened(): raise RuntimeError(f'failed to open video writer: {video_path}')
first=first_image.copy()
cv2.putText(first,f'episode={episode} step=0 success=0',(6,18),cv2.FONT_HERSHEY_SIMPLEX,.42,(0,255,255),1,cv2.LINE_AA)
writer.write(cv2.cvtColor(first,cv2.COLOR_RGB2BGR))
try:
while not done:
current_image=image(obs)
request={'observation/image':current_image,'observation/proprio':proprio(env,args.proprio_convention),
'prompt':env.unwrapped.get_language_instruction(),
**history.request(current_image),
'inference_seed':episode*10000+steps}
await ws.send(msgpack.packb(request,default=pack))
response=await ws.recv()
if isinstance(response,str): raise RuntimeError(response)
chunk=np.asarray(msgpack.unpackb(response,object_hook=unpack)['action'],np.float32)
if chunk.shape!=(32,7) or not np.isfinite(chunk).all(): raise ValueError('invalid Zeva action chunk')
query_step = steps
for chunk_offset, prediction in enumerate(chunk[:args.exec_horizon]):
normalized=np.clip(prediction,-1,1)
raw=0.5*(normalized[:6]+1)*(Q99-Q01)+Q01
command=np.r_[raw[:3],Rotation.from_euler('xyz',raw[3:]).as_rotvec(),1.0 if normalized[6]>0 else -1.0].astype(np.float32)
before=proprio(env,args.proprio_convention)
obs,_,terminated,truncated,info=env.step(command)
after=proprio(env,args.proprio_convention)
reached_raw=reached_delta_raw(before[:6],after[:6],command[6])
history.append_reached_delta(reached_raw,image(obs)); steps+=1
success=success or bool(info.get('success',False))
for key in ('moved_correct_obj','moved_wrong_obj','is_src_obj_grasped','consecutive_grasp','src_on_target'):
progress[key]=progress.get(key,False) or bool(info.get(key,False))
if args.trace_dir:
trace.append(dict(step=steps, query_step=query_step, chunk_offset=chunk_offset,
prediction=prediction.tolist(), command=command.tolist(),
reached_delta_raw=reached_raw.tolist(),
reached_delta_normalized=history.partial[-1].tolist() if history.partial else history.transitions[-1][-1].tolist(),
proprio_before=before.tolist(),proprio=after.tolist(),
info={k:bool(info.get(k,False)) for k in (*progress,'success')}))
done=bool(terminated or truncated)
if writer is not None:
frame=image(obs).copy()
cv2.putText(frame,f'episode={episode} step={steps} success={int(success)}',(6,18),cv2.FONT_HERSHEY_SIMPLEX,.42,(0,255,255),1,cv2.LINE_AA)
cv2.putText(frame,'action='+np.array2string(command,precision=2,suppress_small=True),(6,244),cv2.FONT_HERSHEY_SIMPLEX,.28,(255,255,255),1,cv2.LINE_AA)
writer.write(cv2.cvtColor(frame,cv2.COLOR_RGB2BGR))
if done: break
if args.trace_dir:
trace_dir=Path(args.trace_dir); trace_dir.mkdir(parents=True,exist_ok=True)
(trace_dir/f'episode_{episode:02d}.json').write_text(json.dumps(trace,indent=2))
results.append(dict(seed=episode,steps=steps,success=success,progress=progress,video=str(video_path) if video_path else None))
print(json.dumps(results[-1]),flush=True)
finally:
if writer is not None:
writer.release()
# OpenCV's mp4v output is valid MPEG-4 Part 2 but is not
# playable in some browser previews. Re-mux to broadly
# supported H.264/yuv420p in place.
if video_path is not None:
temp_path=video_path.with_name(video_path.stem+'.h264.tmp.mp4')
try:
subprocess.run(['ffmpeg','-y','-loglevel','error','-i',str(video_path),
'-c:v','libx264','-pix_fmt','yuv420p','-movflags','+faststart',str(temp_path)],
check=True)
temp_path.replace(video_path)
except (FileNotFoundError, subprocess.CalledProcessError):
if temp_path.exists(): temp_path.unlink()
env.close()
report=dict(status='complete',policy_controlled=True,task='widowx_carrot_on_plate',
proprio_convention=args.proprio_convention,protocol_version='widowx_full_v3',
exec_horizon=args.exec_horizon,render_device=args.render_device,memory_mode=args.memory_mode,
checkpoint=args.checkpoint,max_seconds=args.max_seconds,seed_start=args.seed_start,
max_steps=(args.max_steps if args.max_steps is not None else
max(1, int(math.ceil(args.max_seconds * 5.0)))
if args.max_seconds is not None else None),
episodes=results,success_rate=float(np.mean([x['success'] for x in results])))
path=Path(args.output); path.parent.mkdir(parents=True,exist_ok=True)
temp=path.with_suffix('.tmp'); temp.write_text(json.dumps(report,indent=2)); temp.replace(path)
if __name__=='__main__':
p=argparse.ArgumentParser(); p.add_argument('--port',type=int,default=18765)
p.add_argument('--episodes',type=int,default=10); p.add_argument('--checkpoint',required=True); p.add_argument('--output',required=True)
p.add_argument('--seed-start',type=int,default=0,help='first environment and inference seed')
p.add_argument('--memory-mode',choices=('on','bank_off'),default='on')
p.add_argument('--render-device',default='cuda:0',
help='SAPIEN renderer device within CUDA_VISIBLE_DEVICES; default explicitly selects its first GPU')
p.add_argument('--exec-horizon',type=int,default=4,choices=range(1,33),
help='controls to execute per prediction, independent of the 32-step prediction horizon')
p.add_argument('--video-dir',default=None,help='directory for one MP4 per evaluation episode')
p.add_argument('--trace-dir',default=None,help='optional per-step action/state/progress JSON directory')
limit=p.add_mutually_exclusive_group()
limit.add_argument('--max-seconds',type=float,default=None,
help='override episode TimeLimit in seconds at the 5 Hz WidowX control rate')
limit.add_argument('--max-steps',type=int,default=None,
help='maximum environment control steps per episode')
p.add_argument('--proprio-convention',choices=('bridge_tool','legacy_tcp'),default='bridge_tool',help='legacy_tcp is for regression comparison only')
args=p.parse_args()
if args.max_steps is not None and args.max_steps < 1:
p.error('--max-steps must be positive')
if args.seed_start < 0:
p.error('--seed-start must be nonnegative')
asyncio.run(evaluate(args))