Files
VLA/scripts/draw_openloop_plot.py
T

1107 lines
36 KiB
Python
Raw Normal View History

2026-06-15 11:40:00 +08:00
#!/usr/bin/env python3
"""Open-loop evaluation over WebSocket using LeRobot-format datasets.
Loads episodes from a local LeRobot v3 dataset, sends observations to a running
Wall-X websocket server, collects predicted action chunks, and plots them against
ground-truth trajectories.
Reference: ``infer_openloop_websocket.py`` (websocket client + open-loop loop).
For single-arm LIBERO checkpoints, start the server with raw model output::
python -m wall_x._vendor.harrix.serving.launch_serving \\
--env X2ROBOT --port 32194 \\
--no-serialize-actions \\
model-config:server-model-config \\
--model-config.checkpoint-path /path/to/ckpt \\
--model-config.train-config-path /path/to/libero.yml \\
--model-config.action-horizon 10 \\
--model-config.robot-type desktop
"""
from __future__ import annotations
2025-10-24 17:29:12 +08:00
import argparse
2026-06-15 11:40:00 +08:00
import asyncio
import base64
import json
import logging
import sys
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
import numpy as np
import yaml
logger = logging.getLogger(__name__)
2025-09-07 14:59:17 +08:00
2026-06-15 11:40:00 +08:00
def require_cv2():
try:
import cv2
except ImportError as exc:
raise RuntimeError(
"Open-loop image encoding requires opencv-python-headless. "
"Install the repository requirements first."
) from exc
return cv2
2025-09-11 13:18:33 +08:00
2026-06-15 11:40:00 +08:00
def require_msgpack():
try:
import msgpack
import msgpack_numpy as m
except ImportError as exc:
raise RuntimeError(
"Open-loop websocket evaluation requires msgpack-numpy. "
"Install the repository requirements first."
) from exc
m.patch()
return msgpack
def require_websockets():
try:
import websockets
except ImportError as exc:
raise RuntimeError(
"Open-loop websocket evaluation requires websockets. "
"Install the repository requirements first."
) from exc
return websockets
_DEFAULT_CAM_WS_KEYS = {
"observation.images.faceImg": "face_view",
"observation.images.rightImg": "right_wrist_view",
"observation.images.leftImg": "left_wrist_view",
"observation.images.move1Img": "move1_view",
}
@dataclass
class EpisodeArrays:
"""Episode state/action plus lazily decoded camera frames."""
episode_index: int
instruction: str
states: np.ndarray # (T, D_state)
actions: np.ndarray # (T, D_action)
camera_keys: list[str]
frame_offset: int = 0
num_steps: int = 0
_dataset: Any = field(default=None, repr=False)
_image_cache: dict[int, dict[str, np.ndarray]] = field(
default_factory=dict, repr=False
)
def get_frame_images(self, frame_idx: int) -> dict[str, np.ndarray]:
"""Decode camera frames for one observation index (cached)."""
if frame_idx in self._image_cache:
return self._image_cache[frame_idx]
if self._dataset is None:
raise RuntimeError("Episode image loader is not initialized.")
if not self._image_cache:
logger.info(
"Decoding camera frames on demand (first obs_idx=%d)", frame_idx
)
item = self._dataset[frame_idx]
images = {
cam_key: tensor_to_rgb_uint8(item[cam_key])
for cam_key in self.camera_keys
if cam_key in item
}
self._image_cache[frame_idx] = images
return images
def _repo_root() -> Path:
return Path(__file__).resolve().parents[1]
def _add_source_root_if_needed() -> None:
repo_root = _repo_root()
if (repo_root / "wall_x").is_dir() and str(repo_root) not in sys.path:
sys.path.insert(0, str(repo_root))
def load_train_config(path: str | Path) -> dict[str, Any]:
with open(path, "r", encoding="utf-8") as f:
cfg = yaml.load(f, Loader=yaml.FullLoader) or {}
if "data" not in cfg:
cfg["data"] = {}
cfg["data"]["model_type"] = cfg.get("model_type")
return cfg
def parse_int_list(value: str | None) -> list[int] | None:
if value is None or not value.strip():
return None
return [int(x.strip()) for x in value.split(",") if x.strip()]
def _task_block(train_config: dict[str, Any]) -> dict[str, Any]:
return train_config.get("task") or {}
def _dof_config(train_config: dict[str, Any]) -> dict[str, int]:
return (
train_config.get("dof_config")
or _task_block(train_config).get("dof_config")
or {}
)
def _agent_pos_config(train_config: dict[str, Any]) -> dict[str, int]:
return (
train_config.get("agent_pos_config")
or _task_block(train_config).get("agent_pos_config")
or {}
)
def real_vector_dim(config_block: dict[str, int]) -> int:
"""Sum dof/agent dims excluding ``action_padding``."""
return sum(d for k, d in config_block.items() if k != "action_padding" and d > 0)
def slice_by_config(vec: np.ndarray, config_block: dict[str, int]) -> np.ndarray:
"""Keep only non-padding dims from a 1-D vector (state / single action row)."""
vec = np.asarray(vec, dtype=np.float32).reshape(-1)
if not config_block:
return vec
parts: list[np.ndarray] = []
start = 0
for key, dim in config_block.items():
if key == "action_padding":
start += dim
continue
end = start + dim
if end <= vec.shape[0]:
parts.append(vec[start:end])
start = end
if parts:
return np.concatenate(parts, axis=0).astype(np.float32)
return vec
def strip_padding_columns(arr: np.ndarray, config_block: dict[str, int]) -> np.ndarray:
"""Drop ``action_padding`` columns from a [T, D] action chunk."""
arr = np.asarray(arr, dtype=np.float32)
if arr.ndim == 1:
return slice_by_config(arr, config_block)
if not config_block:
return arr
parts: list[np.ndarray] = []
start = 0
for key, dim in config_block.items():
if key == "action_padding":
start += dim
continue
parts.append(arr[:, start : start + dim])
start += dim
if parts:
return np.concatenate(parts, axis=1).astype(np.float32)
return arr
def tensor_to_rgb_uint8(img: Any) -> np.ndarray:
arr = np.asarray(img)
if arr.ndim == 3 and arr.shape[0] in (1, 3) and arr.shape[0] != arr.shape[-1]:
arr = np.transpose(arr, (1, 2, 0))
if np.issubdtype(arr.dtype, np.floating):
arr = np.clip(arr, 0.0, 1.0) * 255.0
return arr.astype(np.uint8)
def encode_image_rgb(image: np.ndarray) -> str:
"""JPEG base64; input must be RGB uint8."""
cv2 = require_cv2()
bgr = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
ok, buffer = cv2.imencode(".jpg", bgr)
if not ok:
raise RuntimeError("cv2.imencode failed")
return base64.b64encode(buffer).decode("utf-8")
def build_camera_ws_mapping(train_config: dict[str, Any]) -> dict[str, str]:
key_mappings = (train_config.get("data") or {}).get("key_mappings") or {}
cam_map = key_mappings.get("camera") or {}
if cam_map:
return dict(cam_map)
return dict(_DEFAULT_CAM_WS_KEYS)
def resolve_feature_keys(
train_config: dict[str, Any],
state_key: str | None,
action_key: str | None,
) -> tuple[str, str]:
key_mappings = (train_config.get("data") or {}).get("key_mappings") or {}
sk = state_key or key_mappings.get("state") or "observation.state"
ak = action_key or key_mappings.get("action") or "action"
return sk, ak
def _arm_follow_pos_from_agent_cfg(
state_vec: np.ndarray,
agent_cfg: dict[str, int],
arm_prefix: str,
) -> list[float] | None:
"""Pack one arm into websocket ``follow{1,2}_pos`` layout: pos3 + rpy3 + grip1."""
from wall_x._vendor.x2robot_utils import geometry as geom
pos = rot = grip = None
start = 0
for key, dim in agent_cfg.items():
if key == "action_padding":
start += int(dim)
continue
end = start + int(dim)
chunk = (
state_vec[start:end]
if end <= state_vec.shape[0]
else np.zeros(int(dim), dtype=np.float32)
)
start = end
if not key.startswith(arm_prefix):
continue
if "cartesian_pos" in key:
pos = chunk.reshape(-1)
elif "rotation_6d" in key.lower():
rot = geom.so3_to_euler_zyx_batch_nb(chunk.reshape(1, -1)).reshape(-1)
elif "rotation" in key:
rot = chunk.reshape(-1)
elif "gripper" in key:
grip = chunk.reshape(-1)
if pos is None or rot is None or grip is None:
return None
packed = np.concatenate([pos[:3], rot[:3], grip[:1]], axis=0).astype(np.float32)
if packed.shape[0] != 7:
return None
return packed.tolist()
def build_state_payload(
state_vec: np.ndarray,
train_config: dict[str, Any],
) -> dict[str, list[float]]:
"""Build websocket ``state`` dict from a flat proprio vector."""
state_vec = np.asarray(state_vec, dtype=np.float32).reshape(-1)
# LeRobot dual-arm export in this workflow is already raw follow-pos 14D:
# [L_pos3, L_rot3, L_grip1, R_pos3, R_rot3, R_grip1].
# Prefer direct passthrough and avoid parsing with agent_pos_config
# (which may be 6D/padded and would mis-slice a 14D state).
if state_vec.shape[0] >= 14:
return {
"follow1_pos": state_vec[:7].tolist(),
"follow2_pos": state_vec[7:14].tolist(),
}
if state_vec.shape[0] >= 7:
return {"follow2_pos": state_vec.tolist()}
agent_cfg = _agent_pos_config(train_config)
if agent_cfg:
has_left = any(k.startswith("follow_left_") for k in agent_cfg)
has_right = any(k.startswith("follow_right_") for k in agent_cfg)
payload: dict[str, list[float]] = {}
if has_left:
left = _arm_follow_pos_from_agent_cfg(state_vec, agent_cfg, "follow_left_")
if left is not None:
payload["follow1_pos"] = left
if has_right:
right = _arm_follow_pos_from_agent_cfg(
state_vec, agent_cfg, "follow_right_"
)
if right is not None:
payload["follow2_pos"] = right
if payload:
return payload
# Legacy single-arm LIBERO path: only follow_right_* in agent_pos_config.
real_dim = real_vector_dim(agent_cfg)
if real_dim > 0 and state_vec.shape[0] >= real_dim:
sliced = slice_by_config(state_vec, agent_cfg)
if sliced.shape[0] == 0:
sliced = state_vec[:real_dim]
right = _arm_follow_pos_from_agent_cfg(
state_vec, agent_cfg, "follow_right_"
)
if right is not None:
return {"follow2_pos": right}
return {"follow2_pos": sliced.tolist()}
raise ValueError(f"Unsupported state dimension: {state_vec.shape[0]}")
def decode_predict_action(
predict_action: Any,
train_config: dict[str, Any],
) -> np.ndarray:
"""Convert raw ``predict_action`` [H, D_model] to flat action rows per step."""
if hasattr(predict_action, "detach"):
predict_action = predict_action.detach().cpu().numpy()
pa = np.asarray(predict_action, dtype=np.float32)
if pa.ndim == 3:
pa = pa[0]
dof_cfg = _dof_config(train_config)
if dof_cfg:
parts: list[np.ndarray] = []
start = 0
for key, dim in dof_cfg.items():
if key == "action_padding":
start += dim
continue
parts.append(pa[:, start : start + dim])
start += dim
if parts:
return np.concatenate(parts, axis=1).astype(np.float32)
try:
from wall_x._vendor.harrix.envs.libero_common import decode_chunk
chunk = decode_chunk(pa, train_config)
if chunk.ndim == 2 and chunk.shape[1] > 0:
return chunk.astype(np.float32)
except Exception:
pass
return pa.astype(np.float32)
def extract_follow_pos_14d_from_response(result: dict[str, Any]) -> np.ndarray:
"""Read serialized ``follow{1,2}_pos`` action rows (skip state row 0) as ``[H, 14]``."""
if "follow2_pos" not in result or "follow1_pos" not in result:
raise KeyError("Response missing follow1_pos/follow2_pos for 14D eval.")
left = np.asarray(result["follow1_pos"], dtype=np.float32)
right = np.asarray(result["follow2_pos"], dtype=np.float32)
if left.ndim == 1:
left = left.reshape(1, -1)
if right.ndim == 1:
right = right.reshape(1, -1)
if left.shape[0] < 2 or right.shape[0] < 2:
raise ValueError(
f"follow_pos response must include state+action rows, got "
f"left={left.shape}, right={right.shape}"
)
horizon = min(left.shape[0] - 1, right.shape[0] - 1)
return np.concatenate([left[1 : 1 + horizon], right[1 : 1 + horizon]], axis=1)
def extract_action_chunk_from_response(
result: dict[str, Any],
train_config: dict[str, Any],
) -> np.ndarray:
if "predict_action" in result:
return decode_predict_action(result["predict_action"], train_config)
if "follow2_pos" in result and "follow1_pos" in result:
return extract_follow_pos_14d_from_response(result)
if "follow2_pos" in result:
right = np.asarray(result["follow2_pos"], dtype=np.float32)
if right.ndim == 1:
right = right.reshape(1, -1)
return right[1:] if right.shape[0] > 1 else right
if "action" in result:
action = np.asarray(result["action"], dtype=np.float32)
if action.ndim == 1:
action = action[np.newaxis, :]
return action
raise KeyError(
"Response has no action fields. For single-arm LIBERO, restart the server with "
"`--no-serialize-actions` so responses include `predict_action`."
)
def _hf_row_to_numpy(row: Any, key: str) -> np.ndarray:
value = row[key]
if hasattr(value, "numpy"):
value = value.numpy()
return np.asarray(value, dtype=np.float32).reshape(-1)
def _episode_instruction(ds: Any, row_index: int = 0) -> str:
row = ds.hf_dataset[row_index]
task_idx = row["task_index"]
if hasattr(task_idx, "item"):
task_idx = task_idx.item()
return str(ds.meta.tasks.iloc[task_idx].name)
def resolve_lerobot_dataset_paths(dataset_root: str | Path) -> tuple[str, Path]:
"""Return ``(repo_id, root)`` for a local LeRobot dataset directory."""
root = Path(dataset_root).expanduser().resolve()
if not root.is_dir():
raise FileNotFoundError(f"LeRobot dataset root not found: {root}")
return root.name, root
def open_local_lerobot_dataset(
dataset_root: str | Path,
*,
episodes: list[int] | None = None,
video_backend: str = "pyav",
) -> Any:
"""Open a local LeRobot v3 dataset without HuggingFace Hub fallback."""
from lerobot.datasets.lerobot_dataset import LeRobotDataset
repo_id, root = resolve_lerobot_dataset_paths(dataset_root)
try:
return LeRobotDataset(
repo_id,
root=root,
episodes=episodes,
video_backend=video_backend,
)
except TypeError:
pass
import torch
from lerobot.datasets.lerobot_dataset import (
CODEBASE_VERSION,
LeRobotDatasetMetadata,
)
from lerobot.datasets.utils import (
get_hf_features_from_features,
hf_transform_to_torch,
load_nested_dataset,
)
from lerobot.datasets.video_utils import get_safe_default_codec
meta = LeRobotDatasetMetadata(repo_id, root=root)
features = get_hf_features_from_features(meta.features)
hf_dataset = load_nested_dataset(
root / "data", features=features, episodes=episodes
)
hf_dataset.set_transform(hf_transform_to_torch)
if episodes is not None:
available = {
int(ep.item()) if hasattr(ep, "item") else int(ep)
for ep in hf_dataset.unique("episode_index")
}
missing = set(episodes) - available
if missing:
raise ValueError(
f"Episodes {sorted(missing)} not found under {root}. "
f"Available in loaded parquet: {sorted(available)[:10]}"
f"{'...' if len(available) > 10 else ''}"
)
if meta.video_keys:
check_eps = (
episodes if episodes is not None else list(range(meta.total_episodes))
)
for ep_idx in check_eps:
for vid_key in meta.video_keys:
video_path = root / meta.get_video_file_path(ep_idx, vid_key)
if not video_path.exists():
raise FileNotFoundError(
"LeRobot camera videos are missing for open-loop eval.\n"
f" dataset root: {root}\n"
f" first missing file: {video_path}\n"
"Parquet state/action may exist, but this script also needs "
"mp4 files under videos/."
)
ds = LeRobotDataset.__new__(LeRobotDataset)
ds.repo_id = repo_id
ds.root = root
ds.image_transforms = None
ds.delta_timestamps = None
ds.episodes = episodes
ds.tolerance_s = 1e-4
ds.revision = CODEBASE_VERSION
ds.video_backend = video_backend or get_safe_default_codec()
ds.delta_indices = None
ds.meta = meta
ds.hf_dataset = hf_dataset
ds._lazy_loading = False
ds._absolute_to_relative_idx = None
if episodes is not None:
ds._absolute_to_relative_idx = {
abs_idx.item() if isinstance(abs_idx, torch.Tensor) else abs_idx: rel_idx
for rel_idx, abs_idx in enumerate(hf_dataset["index"])
}
ds.image_writer = None
ds.episode_buffer = None
ds.writer = None
ds.latest_episode = None
ds._current_file_start_frame = None
ds._streaming_encoder = None
ds.batch_encoding_size = 1
ds.episodes_since_last_encoding = 0
return ds
def load_lerobot_metadata(dataset_root: str | Path) -> Any:
from lerobot.datasets.lerobot_dataset import LeRobotDatasetMetadata
repo_id, root = resolve_lerobot_dataset_paths(dataset_root)
return LeRobotDatasetMetadata(repo_id, root=root)
def plan_eval_frame_range(
num_steps: int,
start_idx: int,
step_stride: int,
action_horizon: int,
max_inferences: int | None,
) -> tuple[int, int, int]:
"""Return ``(tabular_start, tabular_end, expected_video_decodes)``."""
tabular_start = max(0, int(start_idx))
if max_inferences is None:
return tabular_start, num_steps, -1
tabular_end = tabular_start
idx = tabular_start
infer_count = 0
video_decodes = 0
while idx <= num_steps - action_horizon - 1:
if infer_count >= max_inferences:
break
record_n = min(step_stride, action_horizon, num_steps - idx)
if record_n <= 0:
break
tabular_end = max(tabular_end, idx + record_n)
video_decodes += 1
idx += record_n
infer_count += 1
return tabular_start, min(num_steps, tabular_end), video_decodes
def load_episode_arrays(
dataset_root: str | Path,
episode_index: int,
state_key: str,
action_key: str,
camera_keys: list[str],
*,
frame_start: int = 0,
frame_end: int | None = None,
preload_all_images: bool = False,
) -> EpisodeArrays:
ds = open_local_lerobot_dataset(
dataset_root,
episodes=[episode_index],
video_backend="pyav",
)
num_steps = len(ds.hf_dataset)
if num_steps == 0:
raise ValueError(f"Episode {episode_index} is empty under {dataset_root}")
start = max(0, int(frame_start))
end = num_steps if frame_end is None else min(num_steps, int(frame_end))
states = [_hf_row_to_numpy(ds.hf_dataset[i], state_key) for i in range(start, end)]
actions = [
_hf_row_to_numpy(ds.hf_dataset[i], action_key) for i in range(start, end)
]
instruction = _episode_instruction(ds, start)
episode = EpisodeArrays(
episode_index=episode_index,
instruction=instruction,
states=np.stack(states, axis=0),
actions=np.stack(actions, axis=0),
camera_keys=list(camera_keys),
frame_offset=start,
num_steps=num_steps,
_dataset=ds,
)
if preload_all_images:
for frame_idx in range(start, end):
episode.get_frame_images(frame_idx)
return episode
def build_obs_payload(
episode: EpisodeArrays,
frame_idx: int,
train_config: dict[str, Any],
cam_ws_mapping: dict[str, str],
extra_view_keys: list[str],
) -> dict[str, Any]:
local_idx = frame_idx - episode.frame_offset
if local_idx < 0 or local_idx >= episode.states.shape[0]:
raise IndexError(
f"Frame {frame_idx} is outside loaded tabular range "
f"[{episode.frame_offset}, {episode.frame_offset + episode.states.shape[0]})."
)
state = episode.states[local_idx]
frame_images = episode.get_frame_images(frame_idx)
views: dict[str, str] = {}
ref_shape: tuple[int, int, int] | None = None
for cam_key, ws_key in cam_ws_mapping.items():
if cam_key not in frame_images:
continue
rgb = frame_images[cam_key]
ref_shape = rgb.shape
views[ws_key] = encode_image_rgb(rgb)
for ws_key in extra_view_keys:
if ws_key in views:
continue
if ref_shape is None:
ref_shape = (256, 256, 3)
views[ws_key] = encode_image_rgb(np.zeros(ref_shape, dtype=np.uint8))
return {
"state": build_state_payload(state, train_config),
"views": views,
"instruction": episode.instruction,
}
def _dim_label(i: int, dim: int) -> str:
arm7 = [
"pos_x",
"pos_y",
"pos_z",
"rot_x",
"rot_y",
"rot_z",
"gripper",
]
if dim == 14:
prefix = "L_" if i < 7 else "R_"
return prefix + arm7[i % 7]
if dim == 7:
return arm7[i] if i < len(arm7) else f"dim_{i}"
if dim == 12:
prefix = "L_" if i < 6 else "R_"
return f"{prefix}rot6_{i % 6}"
if dim == 10:
names = [
"pos_x",
"pos_y",
"pos_z",
"rot6_0",
"rot6_1",
"rot6_2",
"rot6_3",
"rot6_4",
"rot6_5",
"gripper",
]
return names[i] if i < len(names) else f"dim_{i}"
return f"dim_{i}"
def follow_pos_14d_to_rotation_6d(rows: list | np.ndarray) -> np.ndarray:
"""Convert ``follow1_pos + follow2_pos`` rows from euler rotation to rot6D only."""
from wall_x._vendor.x2robot_utils import geometry as geom
arr = np.asarray(rows, dtype=np.float32)
if arr.ndim == 1:
arr = arr[np.newaxis, :]
if arr.shape[1] < 14:
raise ValueError(f"Expected follow_pos_14d rows, got shape {arr.shape}")
left_rot6 = geom.euler_to_matrix_zyx_6d_nb(arr[:, 3:6]).astype(np.float32)
right_rot6 = geom.euler_to_matrix_zyx_6d_nb(arr[:, 10:13]).astype(np.float32)
return np.concatenate([left_rot6, right_rot6], axis=1).astype(np.float32)
def plot_openloop(
action_pred_list: list | np.ndarray,
action_gt_list: list | np.ndarray,
save_path: Path,
*,
title: str = "",
action_l1: float | None = None,
) -> None:
"""Plot aligned GT vs predicted action trajectories."""
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
gt_rows = np.asarray(action_gt_list, dtype=np.float32)
pred_rows = np.asarray(action_pred_list, dtype=np.float32)
assert (
gt_rows.shape == pred_rows.shape
), f"shape mismatch: pred={pred_rows.shape} gt={gt_rows.shape}"
n, dim = gt_rows.shape
fig, axes = plt.subplots(dim, 1, figsize=(12, 3.5 * dim), sharex=True)
if dim == 1:
axes = [axes]
for i, ax in enumerate(axes):
ax.plot(gt_rows[:, i], label="Ground Truth", color="blue", linewidth=1.5)
ax.plot(pred_rows[:, i], label="Model Output", color="orange", linewidth=1.5)
ax.set_title(_dim_label(i, dim))
ax.set_ylabel("Action Value")
ax.grid(True, alpha=0.25)
ax.legend(loc="best")
axes[-1].set_xlabel("Time Step")
axes[0].set_xticks(np.arange(0, n, step=max(1, min(10, max(n // 15, 1)))))
l1_note = f"Mean L1: {action_l1:.6f}" if action_l1 is not None else ""
if title and l1_note:
fig.suptitle(f"{title}\n{l1_note}")
elif title:
fig.suptitle(title)
elif l1_note:
fig.suptitle(l1_note)
fig.tight_layout(rect=[0, 0, 1, 0.97])
save_path.parent.mkdir(parents=True, exist_ok=True)
out = save_path.with_suffix(".jpg") if save_path.suffix != ".jpg" else save_path
fig.savefig(out, dpi=200)
plt.close(fig)
logger.info("Saved plot -> %s", out)
def compute_second_diff(arr: np.ndarray) -> np.ndarray:
d1 = arr[1:] - arr[:-1]
return d1[1:] - d1[:-1]
async def run_openloop_eval(
uri: str,
dataset_root: str,
train_config_path: str,
save_dir: str,
episode_indices: list[int] | None,
start_ratio: float,
stride: int | None,
max_inferences: int | None,
state_key: str | None,
action_key: str | None,
extra_view_keys: list[str],
preload_all_images: bool = False,
plot_rotation_6d: bool = False,
) -> None:
train_config = load_train_config(train_config_path)
state_key, action_key = resolve_feature_keys(train_config, state_key, action_key)
cam_ws_mapping = build_camera_ws_mapping(train_config)
camera_keys = list(cam_ws_mapping.keys())
action_horizon = (
train_config.get("action_horizon_flow")
or _task_block(train_config).get("action_horizon_flow")
or _task_block(train_config).get("action_horizon")
or train_config.get("action_horizon")
or 10
)
step_stride = stride if stride is not None else int(action_horizon)
meta = load_lerobot_metadata(dataset_root)
if episode_indices is None:
episode_indices = [0]
episode_indices = [i for i in episode_indices if 0 <= i < meta.total_episodes]
if not episode_indices:
raise ValueError("No valid episode indices to evaluate.")
save_root = Path(save_dir)
save_root.mkdir(parents=True, exist_ok=True)
msgpack = require_msgpack()
websockets = require_websockets()
async with websockets.connect(
uri,
ping_interval=None,
ping_timeout=None,
max_size=None,
) as websocket:
metadata = msgpack.unpackb(await websocket.recv())
logger.info("Connected to %s, server metadata: %s", uri, metadata)
for ep_idx in episode_indices:
ep_meta = meta.episodes[ep_idx]
num_steps = int(
ep_meta.get("length")
if isinstance(ep_meta, dict)
else getattr(ep_meta, "length", 0)
)
if num_steps <= 0:
raise ValueError(f"Episode {ep_idx} has invalid length metadata.")
start_idx = int(start_ratio * num_steps)
tabular_start, tabular_end, expected_decodes = plan_eval_frame_range(
num_steps,
start_idx,
step_stride,
int(action_horizon),
max_inferences,
)
logger.info(
"Loading episode %d tabular frames [%d, %d) / %d "
"(images on demand%s)...",
ep_idx,
tabular_start,
tabular_end,
num_steps,
f", ~{expected_decodes} obs frames" if expected_decodes >= 0 else "",
)
episode = load_episode_arrays(
dataset_root,
ep_idx,
state_key=state_key,
action_key=action_key,
camera_keys=camera_keys,
frame_start=tabular_start,
frame_end=tabular_end,
preload_all_images=preload_all_images,
)
logger.info(
"Episode %d tabular ready: %d frames loaded, video decode per inference",
ep_idx,
episode.states.shape[0],
)
gt_full = np.asarray(episode.actions, dtype=np.float32)
if gt_full.ndim == 1:
gt_full = gt_full[np.newaxis, :]
pred_full = np.full_like(gt_full, np.nan)
aligned_gt: list[list[float]] = []
aligned_pred: list[list[float]] = []
obs_infer_points: list[tuple[int, int]] = []
idx = start_idx
infer_count = 0
while idx <= num_steps - action_horizon - 1:
if max_inferences is not None and infer_count >= max_inferences:
break
payload = build_obs_payload(
episode,
idx,
train_config,
cam_ws_mapping,
extra_view_keys=extra_view_keys,
)
await websocket.send(msgpack.packb(payload, use_bin_type=True))
raw = await websocket.recv()
if isinstance(raw, str):
hint = ""
if "get_serialized_actions" in raw:
hint = (
"\n\nHint: restart the server with `--no-serialize-actions` "
"so responses include `predict_action`."
)
raise RuntimeError(f"Server error at frame {idx}:\n{raw}{hint}")
result = msgpack.unpackb(raw, raw=False)
pred_chunk = extract_action_chunk_from_response(result, train_config)
if pred_chunk.ndim == 1:
pred_chunk = pred_chunk[np.newaxis, :]
if pred_chunk.shape[1] != gt_full.shape[1]:
raise ValueError(
f"Prediction action dim {pred_chunk.shape[1]} does not match "
f"dataset action dim {gt_full.shape[1]}."
)
record_n = min(
step_stride, action_horizon, num_steps - idx, pred_chunk.shape[0]
)
if record_n <= 0:
break
local_idx = idx - episode.frame_offset
gt_chunk = gt_full[local_idx : local_idx + record_n]
pred_chunk = pred_chunk[:record_n]
record_n = min(record_n, pred_chunk.shape[0], gt_chunk.shape[0])
if record_n <= 0:
break
row0 = len(aligned_gt)
obs_infer_points.append((row0, idx))
pred_full[local_idx : local_idx + record_n] = pred_chunk[:record_n]
aligned_gt.extend(gt_chunk[:record_n].tolist())
aligned_pred.extend(pred_chunk[:record_n].tolist())
logger.info(
"episode=%d obs_idx=%d row0=%d record_n=%d pred_shape=%s gt_dim=%d",
ep_idx,
idx,
row0,
record_n,
pred_chunk.shape,
gt_chunk.shape[1],
)
idx += record_n
infer_count += 1
if not obs_infer_points:
logger.warning("Episode %d: no inference steps executed.", ep_idx)
continue
ep_tag = f"ep{ep_idx}"
valid_mask = ~np.isnan(pred_full).any(axis=1)
action_l1 = float(
np.mean(np.abs(pred_full[valid_mask] - gt_full[valid_mask]))
)
summary = {
"episode_index": ep_idx,
"instruction": episode.instruction,
"num_steps": num_steps,
"obs_infer_points": obs_infer_points,
"action_horizon": action_horizon,
"stride": step_stride,
"mean_l1": action_l1,
"state_key": state_key,
"action_key": action_key,
}
summary_path = save_root / f"{ep_tag}_summary.json"
summary_path.write_text(
json.dumps(summary, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
logger.info("Episode %d mean L1 = %.6f", ep_idx, action_l1)
ml = min(len(aligned_gt), len(aligned_pred))
if ml < 3:
logger.warning(
"Episode %d: too few aligned rows (%d) to plot.", ep_idx, ml
)
continue
gt_ml = np.asarray(aligned_gt[:ml], dtype=np.float32)
pred_ml = np.asarray(aligned_pred[:ml], dtype=np.float32)
title = f"Episode {ep_idx}: {episode.instruction[:80]}"
plot_openloop(
pred_ml,
gt_ml,
save_root / ep_tag,
title=title,
action_l1=action_l1,
)
if plot_rotation_6d:
gt_rot6d = follow_pos_14d_to_rotation_6d(gt_ml)
pred_rot6d = follow_pos_14d_to_rotation_6d(pred_ml)
rot6d_l1 = float(np.mean(np.abs(pred_rot6d - gt_rot6d)))
plot_openloop(
pred_rot6d,
gt_rot6d,
save_root / f"{ep_tag}_rot6d",
title=f"{title} (rotation 6D)",
action_l1=rot6d_l1,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Open-loop websocket inference on LeRobot datasets.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument("--uri", default="ws://127.0.0.1:32194")
parser.add_argument(
"--dataset-root",
required=True,
help="Local LeRobot dataset root.",
)
parser.add_argument(
"--train-config",
required=True,
help="Training YAML for key mappings / dof layout.",
)
parser.add_argument(
"--save-dir",
default="./openloop_lerobot_plots",
)
parser.add_argument(
"--episode-indices",
default="0",
help="Comma-separated episode indices, e.g. 0,1,2",
)
parser.add_argument(
"--start-ratio",
type=float,
default=0.0,
help="Start open-loop from this fraction of the episode length.",
)
parser.add_argument(
"--stride",
type=int,
default=None,
help="Frames to advance between inferences (default: action_horizon).",
)
parser.add_argument(
"--max-inferences",
type=int,
default=None,
help=(
"Cap inference requests per episode; also limits tabular/video loading "
"to the evaluated frame range."
),
)
parser.add_argument(
"--preload-all-images",
action="store_true",
help="Decode all camera frames up front (slow; old behavior).",
)
parser.add_argument(
"--plot-rotation-6d",
action="store_true",
help="Also save a rotation-6D-only plot as ep*_rot6d.jpg.",
)
parser.add_argument(
"--state-key",
default=None,
help="LeRobot state feature key (default: from train config key_mappings).",
)
parser.add_argument(
"--action-key",
default=None,
help="LeRobot action feature key (default: from train config key_mappings).",
)
parser.add_argument(
"--extra-view-keys",
default="left_wrist_view",
help="Websocket view keys to fill with black images when missing from data.",
)
parser.add_argument("--log-level", default="INFO")
return parser.parse_args()
def main() -> int:
args = parse_args()
logging.basicConfig(
level=getattr(logging, args.log_level.upper(), logging.INFO),
format="%(asctime)s - %(levelname)s - %(message)s",
)
_add_source_root_if_needed()
extra_view_keys = [k.strip() for k in args.extra_view_keys.split(",") if k.strip()]
episode_indices = parse_int_list(args.episode_indices)
try:
asyncio.run(
run_openloop_eval(
uri=args.uri,
dataset_root=args.dataset_root,
train_config_path=args.train_config,
save_dir=args.save_dir,
episode_indices=episode_indices,
start_ratio=args.start_ratio,
stride=args.stride,
max_inferences=args.max_inferences,
state_key=args.state_key,
action_key=args.action_key,
extra_view_keys=extra_view_keys,
preload_all_images=args.preload_all_images,
plot_rotation_6d=args.plot_rotation_6d,
)
)
except KeyboardInterrupt:
logger.info("Stopped by user.")
return 130
return 0
2025-09-07 14:59:17 +08:00
2025-09-11 13:18:33 +08:00
2025-10-24 17:29:12 +08:00
if __name__ == "__main__":
2026-06-15 11:40:00 +08:00
raise SystemExit(main())