import gc import json import os import time from pathlib import Path import av import numpy as np import pandas as pd REPO_ROOT = Path(__file__).resolve().parents[2] DEFAULT_OUTPUT = ( REPO_ROOT / "evaluation" / "outputs" / "continuous_6tasks" ) OUTPUT = Path( os.environ.get("CONT_EVAL_OUTPUT", str(DEFAULT_OUTPUT)) ).expanduser().resolve() OUTPUT.mkdir(parents=True, exist_ok=True) os.environ.setdefault( "SMOKE_OUTPUT", str(OUTPUT / "_smoke_import"), ) import offline_replay_smoke as smoke EPISODES = (50, 150, 250, 350, 450, 550) EXECUTE_STEPS = 5 ACTION_HORIZON = 10 RANDOM_SEED = 20260916 ARM_DIMS = smoke.ARM_DIMS GRIPPER_DIMS = smoke.GRIPPER_DIMS JOINT_NAMES = ( "left_joint_1", "left_joint_2", "left_joint_3", "left_joint_4", "left_joint_5", "left_joint_6", "left_gripper", "right_joint_1", "right_joint_2", "right_joint_3", "right_joint_4", "right_joint_5", "right_joint_6", "right_gripper", ) def create_policy(): config = smoke.train_config.get_config( smoke.CONFIG_NAME ) assets = smoke.replace_dataclass_fields( config.data.assets, assets_dir=str(smoke.CKPT / "assets"), asset_id="real_piper_x_6tasks", ) data = smoke.replace_dataclass_fields( config.data, repo_id=str(smoke.DATASET), assets=assets, ) config = smoke.replace_dataclass_fields( config, data=data, ) start = time.monotonic() policy = smoke.policy_config.create_trained_policy( config, smoke.CKPT, ) load_seconds = time.monotonic() - start print(f"policy_load_seconds={load_seconds:.3f}") return config, policy, load_seconds def decode_targets_fast(video_path, targets): targets = [float(x) for x in targets] images = [] actual_timestamps = [] errors = [] with av.open(str(video_path)) as container: stream = container.streams.video[0] time_base = float(stream.time_base) seek_time = max( 0.0, targets[0] - 2.0 / smoke.FPS, ) container.seek( int(seek_time / time_base), stream=stream, any_frame=False, backward=True, ) target_index = 0 previous_frame = None previous_timestamp = None for frame in container.decode(stream): if frame.pts is None: continue timestamp = float( frame.pts * stream.time_base ) while ( target_index < len(targets) and timestamp >= targets[target_index] ): target = targets[target_index] if ( previous_frame is not None and abs(previous_timestamp - target) <= abs(timestamp - target) ): selected_frame = previous_frame selected_timestamp = previous_timestamp else: selected_frame = frame selected_timestamp = timestamp error = abs(selected_timestamp - target) if error > 2.0 / smoke.FPS: raise RuntimeError( f"Video timestamp error: " f"target={target}, " f"actual={selected_timestamp}, " f"error={error}, " f"path={video_path}" ) images.append( smoke.chw_uint8( selected_frame.to_ndarray( format="rgb24" ) ) ) actual_timestamps.append( selected_timestamp ) errors.append(error) target_index += 1 if target_index >= len(targets): break previous_frame = frame previous_timestamp = timestamp if len(images) != len(targets): raise RuntimeError( f"Decoded {len(images)}/{len(targets)} " f"targets from {video_path}" ) return images, actual_timestamps, errors def metric(prediction, target, hold, steps): prediction = prediction[:steps] target = target[:steps] hold = hold[:steps] result = {} for name, dims in ( ("all", None), ("arm", ARM_DIMS), ("gripper", GRIPPER_DIMS), ): if dims is None: pred_diff = prediction - target hold_diff = hold - target else: pred_diff = ( prediction[:, dims] - target[:, dims] ) hold_diff = hold[:, dims] - target[:, dims] result[f"model_mae_{name}"] = float( np.mean(np.abs(pred_diff)) ) result[f"hold_mae_{name}"] = float( np.mean(np.abs(hold_diff)) ) result["model_beats_hold"] = bool( result["model_mae_all"] < result["hold_mae_all"] ) return result def add_metrics(row, prefix, values): for key, value in values.items(): row[f"{prefix}_{key}"] = value def main(): config, policy, load_seconds = create_policy() rng = np.random.default_rng(RANDOM_SEED) records = [] overlap_records = [] predictions = [] targets_all = [] executed_rows = [] global_call_index = 0 for task_id, episode_index in enumerate(EPISODES): smoke.EPISODE_INDEX = episode_index metadata = smoke.load_episode_metadata() episode = smoke.load_episode_data(metadata) states = episode["states"] actions = episode["actions"] episode_length = len(states) prompt = str(metadata["tasks"][0]) call_frames = list( range( 0, episode_length - ACTION_HORIZON + 1, EXECUTE_STEPS, ) ) print( f"===== TASK {task_id} " f"EPISODE {episode_index} =====" ) print(f"prompt={prompt}") print(f"length={episode_length}") print(f"num_calls={len(call_frames)}") images_by_camera = {} decode_debug = {} for camera in smoke.CAMERAS: video_path, from_timestamp, _ = ( smoke.video_path_and_range( metadata, camera, ) ) timestamps = [ from_timestamp + frame / smoke.FPS for frame in call_frames ] images, actual, errors = ( decode_targets_fast( video_path, timestamps, ) ) images_by_camera[camera] = images decode_debug[camera] = { "max_error": float(max(errors)), "mean_error": float( np.mean(errors) ), } print(f"video_decode={decode_debug}") policy._model.reset_memory(policy.memory) previous_prediction = None previous_target = None previous_record = None for call_index, frame_index in enumerate( call_frames ): observation = { "state": states[frame_index].astype( np.float32, copy=False, ), "images": { camera: images_by_camera[camera][ call_index ] for camera in smoke.CAMERAS }, "prompt": prompt, "step": call_index, } noise = rng.standard_normal( ( ACTION_HORIZON, int(config.model.action_dim), ) ).astype(np.float32) start = time.monotonic() output = policy.infer( observation, noise=noise, ) wall_ms = ( time.monotonic() - start ) * 1000.0 prediction = np.asarray( output["actions"], dtype=np.float32, ) target = actions[ frame_index : frame_index + ACTION_HORIZON ].astype(np.float32) hold = np.repeat( states[frame_index][None], ACTION_HORIZON, axis=0, ).astype(np.float32) if prediction.shape != (10, 14): raise RuntimeError( f"Bad prediction shape: " f"{prediction.shape}" ) if not np.isfinite(prediction).all(): raise RuntimeError( f"Non-finite output at " f"episode={episode_index}, " f"frame={frame_index}" ) exec_metrics = metric( prediction, target, hold, EXECUTE_STEPS, ) full_metrics = metric( prediction, target, hold, ACTION_HORIZON, ) row = { "global_call_index": global_call_index, "task_id": task_id, "task": prompt, "episode_index": episode_index, "episode_call_index": call_index, "frame_index": frame_index, "memory_phase": call_index % 3, "wall_ms": wall_ms, "policy_infer_ms": float( output["policy_timing"]["infer_ms"] ), } add_metrics(row, "exec5", exec_metrics) add_metrics(row, "full10", full_metrics) records.append(row) predictions.append(prediction) targets_all.append(target) for horizon_index in range(EXECUTE_STEPS): for joint_index, joint_name in enumerate( JOINT_NAMES ): pred_value = float( prediction[ horizon_index, joint_index, ] ) gt_value = float( target[ horizon_index, joint_index, ] ) executed_rows.append( { "task_id": task_id, "task": prompt, "episode_index": episode_index, "call_index": call_index, "source_frame": frame_index, "executed_frame": ( frame_index + horizon_index ), "memory_phase": call_index % 3, "horizon_index": horizon_index, "joint_index": joint_index, "joint_name": joint_name, "joint_group": ( "gripper" if joint_index in (6, 13) else "arm" ), "prediction": pred_value, "ground_truth": gt_value, "signed_error": ( pred_value - gt_value ), "absolute_error": abs( pred_value - gt_value ), } ) if previous_prediction is not None: gt_overlap = ( target[:5] - previous_target[5:10] ) correct_overlap = ( prediction[:5] - previous_prediction[5:10] ) repeat_head = ( prediction[:5] - previous_prediction[:5] ) overlap_mae = float( np.mean(np.abs(correct_overlap)) ) repeat_mae = float( np.mean(np.abs(repeat_head)) ) overlap_records.append( { "task_id": task_id, "task": prompt, "episode_index": episode_index, "previous_call_index": ( call_index - 1 ), "current_call_index": call_index, "previous_frame": frame_index - 5, "current_frame": frame_index, "from_memory_phase": ( (call_index - 1) % 3 ), "to_memory_phase": ( call_index % 3 ), "gt_overlap_max_abs": float( np.max(np.abs(gt_overlap)) ), "correct_overlap_mae": overlap_mae, "repeat_head_mae": repeat_mae, "replay_margin": ( repeat_mae - overlap_mae ), "correct_alignment": bool( overlap_mae < repeat_mae ), "gt_head_tail_motion": float( np.mean( np.abs( previous_target[5:10] - previous_target[:5] ) ) ), "boundary_delta_error": float( np.mean( np.abs( ( prediction[0] - previous_prediction[4] ) - ( target[0] - previous_target[4] ) ) ) ), } ) previous_prediction = prediction previous_target = target previous_record = row global_call_index += 1 if ( call_index % 25 == 0 or call_index + 1 == len(call_frames) ): print( f"call={call_index:04d}/" f"{len(call_frames) - 1:04d} " f"frame={frame_index:04d} " f"phase={call_index % 3} " f"exec5_mae=" f"{exec_metrics['model_mae_all']:.6f} " f"wall_ms={wall_ms:.1f}", flush=True, ) del images_by_camera gc.collect() records_frame = pd.DataFrame(records) overlap_frame = pd.DataFrame(overlap_records) executed_frame = pd.DataFrame(executed_rows) prediction_array = np.stack(predictions) target_array = np.stack(targets_all) records_frame.to_csv( OUTPUT / "continuous_records.csv", index=False, ) overlap_frame.to_csv( OUTPUT / "continuous_overlap.csv", index=False, ) executed_frame.to_csv( OUTPUT / "executed_trajectory_long.csv", index=False, ) np.savez_compressed( OUTPUT / "continuous_predictions.npz", predictions=prediction_array, ground_truth=target_array, task_ids=records_frame[ "task_id" ].to_numpy(dtype=np.int64), episode_indices=records_frame[ "episode_index" ].to_numpy(dtype=np.int64), frame_indices=records_frame[ "frame_index" ].to_numpy(dtype=np.int64), memory_phases=records_frame[ "memory_phase" ].to_numpy(dtype=np.int64), ) per_task_rows = [] for task_id, group in records_frame.groupby( "task_id" ): overlaps = overlap_frame[ overlap_frame["task_id"] == task_id ] per_task_rows.append( { "task_id": int(task_id), "episode_index": int( group["episode_index"].iloc[0] ), "task": group["task"].iloc[0], "num_calls": int(len(group)), "exec5_model_mae_all": float( group[ "exec5_model_mae_all" ].mean() ), "exec5_hold_mae_all": float( group[ "exec5_hold_mae_all" ].mean() ), "exec5_model_mae_arm": float( group[ "exec5_model_mae_arm" ].mean() ), "exec5_model_mae_gripper": float( group[ "exec5_model_mae_gripper" ].mean() ), "exec5_win_rate": float( group[ "exec5_model_beats_hold" ].mean() ), "full10_model_mae_all": float( group[ "full10_model_mae_all" ].mean() ), "full10_hold_mae_all": float( group[ "full10_hold_mae_all" ].mean() ), "correct_alignment_rate": float( overlaps[ "correct_alignment" ].mean() ), "mean_replay_margin": float( overlaps[ "replay_margin" ].mean() ), "mean_boundary_delta_error": float( overlaps[ "boundary_delta_error" ].mean() ), } ) per_task = pd.DataFrame(per_task_rows) per_task.to_csv( OUTPUT / "continuous_per_task.csv", index=False, ) motion_threshold = float( overlap_frame[ "gt_head_tail_motion" ].quantile(0.75) ) high_motion = overlap_frame[ overlap_frame["gt_head_tail_motion"] >= motion_threshold ] warm = records_frame[ records_frame["global_call_index"] >= 3 ] summary = { "evaluation_type": ( "continuous teacher-forced full-episode replay" ), "closed_loop": False, "checkpoint": str(smoke.CKPT), "episodes": list(EPISODES), "num_tasks": len(EPISODES), "num_calls": int(len(records_frame)), "num_executed_actions": int( len(records_frame) * EXECUTE_STEPS ), "execute_steps": EXECUTE_STEPS, "action_horizon": ACTION_HORIZON, "policy_load_seconds": load_seconds, "overall": { "exec5_model_mae_all": float( records_frame[ "exec5_model_mae_all" ].mean() ), "exec5_hold_mae_all": float( records_frame[ "exec5_hold_mae_all" ].mean() ), "exec5_model_mae_arm": float( records_frame[ "exec5_model_mae_arm" ].mean() ), "exec5_model_mae_gripper": float( records_frame[ "exec5_model_mae_gripper" ].mean() ), "exec5_win_rate": float( records_frame[ "exec5_model_beats_hold" ].mean() ), "full10_model_mae_all": float( records_frame[ "full10_model_mae_all" ].mean() ), "full10_hold_mae_all": float( records_frame[ "full10_hold_mae_all" ].mean() ), }, "continuity": { "gt_overlap_max_abs": float( overlap_frame[ "gt_overlap_max_abs" ].max() ), "correct_overlap_mae": float( overlap_frame[ "correct_overlap_mae" ].mean() ), "repeat_head_mae": float( overlap_frame[ "repeat_head_mae" ].mean() ), "mean_replay_margin": float( overlap_frame[ "replay_margin" ].mean() ), "correct_alignment_rate": float( overlap_frame[ "correct_alignment" ].mean() ), "mean_boundary_delta_error": float( overlap_frame[ "boundary_delta_error" ].mean() ), }, "high_motion_continuity": { "threshold": motion_threshold, "num_boundaries": int( len(high_motion) ), "mean_replay_margin": float( high_motion[ "replay_margin" ].mean() ), "correct_alignment_rate": float( high_motion[ "correct_alignment" ].mean() ), }, "timing": { "warm_calls": int(len(warm)), "wall_ms_mean": float( warm["wall_ms"].mean() ), "wall_ms_median": float( warm["wall_ms"].median() ), "wall_ms_p95": float( warm["wall_ms"].quantile(0.95) ), "policy_infer_ms_median": float( warm[ "policy_infer_ms" ].median() ), }, } (OUTPUT / "continuous_summary.json").write_text( json.dumps( summary, indent=2, ensure_ascii=False, ), encoding="utf-8", ) print("===== CONTINUOUS SUMMARY =====") print( json.dumps( summary, indent=2, ensure_ascii=False, ) ) print("===== CONTINUOUS PER TASK =====") print(per_task.to_string(index=False)) print("CONTINUOUS_REPLAY_COMPLETE") if __name__ == "__main__": main()