File size: 11,772 Bytes
68eed61
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
"""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))