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