Download scripts/eval_widowx_simpler.py from DanTim05/DynaTTT: direct link, hf CLI and curl.
- Browser
- Download file 11.8 kB
-
https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/eval_widowx_simpler.py
- Command line
-
hf download hf://spaces/DanTim05/DynaTTT/scripts/eval_widowx_simpler.py
-
curl -L -o eval_widowx_simpler.py https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/eval_widowx_simpler.py
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)) | |