StreamPIReal6_3w / evaluation /scripts /offline_continuous_replay.py
Dengliming's picture
Upload folder using huggingface_hub
00c55c8 verified
Raw
History Blame Contribute Delete
23.5 kB
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()