"""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))