diff --git a/xr_rm_teleop/test/test_act_episode_recorder.py b/xr_rm_teleop/test/test_act_episode_recorder.py index cd9ba81..4913871 100644 --- a/xr_rm_teleop/test/test_act_episode_recorder.py +++ b/xr_rm_teleop/test/test_act_episode_recorder.py @@ -1,14 +1,26 @@ import numpy as np import pytest +try: + import h5py +except ImportError: + h5py = None + from xr_rm_teleop.act_episode_recorder import ( NO_ACTION, ButtonTracker, CameraBuffer, CameraFrame, + EpisodeMetadata, + EpisodeStore, QualityError, RecordingSession, RecordingState, + TaskDirectoryLock, + discard_partial, + next_episode_index, + publish_without_overwrite, + recover_partial_files, select_camera_pair, select_frame, ) @@ -273,3 +285,195 @@ def test_camera_buffer_is_bounded_and_counts_dropped_frames(): assert stats.dropped_frames == 1 assert stats.drop_ratio == pytest.approx(1.0 / 6.0) assert stats.fps == pytest.approx(800.0) + + +requires_h5py = pytest.mark.skipif( + h5py is None, + reason="h5py is not installed", +) + + +def _metadata(): + return EpisodeMetadata( + joint_names=tuple(f"joint_{index}" for index in range(1, 8)) + + ("gripper",), + joint_lower_limits=np.full(7, -3.0, dtype=np.float64), + joint_upper_limits=np.full(7, 3.0, dtype=np.float64), + ) + + +def _episode_sample(seq): + control_ns = seq * 11_111_111 + qpos = np.zeros(8, dtype=np.float32) + qpos[-1] = 1.0 + action = qpos.copy() + return { + "observations/qpos": qpos, + "action": action, + "observations/images/cam_high": _image(seq % 255), + "observations/images/cam_right_wrist": _image((seq + 1) % 255), + "debug/timestamps/control_monotonic_ns": control_ns, + "debug/timestamps/feedback_monotonic_ns": control_ns - 1_000_000, + "debug/timestamps/action_monotonic_ns": control_ns, + "debug/timestamps/cam_high_host_monotonic_ns": control_ns - 2_000_000, + "debug/timestamps/cam_wrist_host_monotonic_ns": control_ns - 3_000_000, + "debug/timestamps/cam_high_hardware_ms": float(seq), + "debug/timestamps/cam_wrist_hardware_ms": float(seq), + "debug/timestamps/cam_high_age_ms": 2.0, + "debug/timestamps/cam_wrist_age_ms": 3.0, + "debug/timestamps/inter_camera_skew_ms": 1.0, + "debug/cameras/cam_high_frame_number": seq, + "debug/cameras/cam_wrist_frame_number": seq, + "debug/control/control_seq": seq, + "debug/control/teleop_active": 1, + "debug/control/action_valid": 1, + "debug/control/command_sent": 1, + "debug/control/target_clamped": 0, + "debug/control/control_fault": 0, + "debug/qp/raw_target": np.zeros(7, dtype=np.float32), + "debug/qp/attempted": 1, + "debug/qp/success": 1, + "debug/qp/duration_ms": 1.0, + "debug/tcp/current_pose": np.zeros(7, dtype=np.float32), + "debug/tcp/raw_target_pose": np.zeros(7, dtype=np.float32), + "debug/tcp/final_target_pose": np.zeros(7, dtype=np.float32), + "debug/tcp/command_velocity": np.zeros(6, dtype=np.float32), + "debug/pico/right_pose": np.zeros(7, dtype=np.float32), + "debug/pico/right_inputs": np.zeros(6, dtype=np.float32), + "debug/pico/left_secondary": 0, + "debug/gripper/target_open": 1, + "debug/gripper/state_open": 1, + "debug/gripper/command_pending": 0, + "debug/gripper/command_failed": 0, + } + + +@requires_h5py +def test_episode_store_writes_act_core_schema(tmp_path): + store = EpisodeStore.create( + tmp_path / "episode_0.partial.hdf5", + _metadata(), + ) + for seq in (100, 103, 106): + store.append(_episode_sample(seq)) + store.close() + + with h5py.File(store.path, "r") as root: + assert root.attrs["sim"] == np.bool_(False) + assert root.attrs["action_alignment"] == "same_step_causal" + assert root["observations/qpos"].shape == (3, 8) + assert root["observations/qpos"].dtype == np.float32 + assert root["action"].shape == (3, 8) + assert root["action"].dtype == np.float32 + assert root["observations/images/cam_high"].shape == ( + 3, + 480, + 640, + 3, + ) + assert root["observations/images/cam_high"].dtype == np.uint8 + assert root["observations/images/cam_right_wrist"].shape == ( + 3, + 480, + 640, + 3, + ) + assert "observations/qvel" not in root + assert "observations/effort" not in root + assert "compress_len" not in root + + +@requires_h5py +def test_episode_store_truncates_every_time_axis_dataset(tmp_path): + store = EpisodeStore.create( + tmp_path / "episode_0.partial.hdf5", + _metadata(), + ) + for seq in (100, 103, 106): + store.append(_episode_sample(seq)) + + store.truncate(2) + store.close() + + lengths = [] + with h5py.File(store.path, "r") as root: + root.visititems( + lambda _name, item: lengths.append(item.shape[0]) + if isinstance(item, h5py.Dataset) + else None + ) + assert lengths + assert set(lengths) == {2} + + +def test_next_index_uses_max_saved_episode_and_ignores_rejected(tmp_path): + (tmp_path / "episode_2.hdf5").touch() + (tmp_path / "episode_9.hdf5").touch() + rejected = tmp_path / "rejected" + rejected.mkdir() + (rejected / "episode_20_bad_20260810.hdf5").touch() + + assert next_episode_index(tmp_path) == 10 + + +def test_publish_never_overwrites_existing_episode(tmp_path): + partial = tmp_path / "episode_1.partial.hdf5" + partial.write_bytes(b"new") + final = tmp_path / "episode_1.hdf5" + final.write_bytes(b"old") + + with pytest.raises(FileExistsError): + publish_without_overwrite(partial, final) + + assert final.read_bytes() == b"old" + assert partial.read_bytes() == b"new" + + +@requires_h5py +def test_recover_marks_readable_and_preserves_unreadable_partial(tmp_path): + readable = tmp_path / "episode_1.partial.hdf5" + store = EpisodeStore.create(readable, _metadata()) + store.append(_episode_sample(100)) + store.close() + unreadable = tmp_path / "episode_2.partial.hdf5" + unreadable.write_bytes(b"not hdf5") + + recovered = recover_partial_files(tmp_path, timestamp="20260810T120000") + + assert len(recovered) == 2 + rejected = tmp_path / "rejected" / ( + "episode_1_crash_recovered_20260810T120000.hdf5" + ) + assert rejected in recovered + with h5py.File(rejected, "r") as root: + assert root.attrs["episode_status"] == "rejected" + assert root.attrs["reject_reason"] == "crash_recovered" + assert root.attrs["interrupted"] == np.bool_(True) + assert not readable.exists() + assert any(path.name.endswith(".partial.hdf5") for path in recovered) + assert unreadable not in recovered + assert not unreadable.exists() + + +def test_discard_partial_removes_only_current_file(tmp_path): + current = tmp_path / "episode_1.partial.hdf5" + current.write_bytes(b"current") + saved = tmp_path / "episode_0.hdf5" + saved.write_bytes(b"saved") + + discard_partial(current) + + assert not current.exists() + assert saved.read_bytes() == b"saved" + assert next_episode_index(tmp_path) == 1 + + +def test_task_directory_lock_rejects_second_recorder(tmp_path): + first = TaskDirectoryLock(tmp_path) + second = TaskDirectoryLock(tmp_path) + first.acquire() + try: + with pytest.raises(BlockingIOError): + second.acquire() + finally: + first.release() diff --git a/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py b/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py index caffc2f..b34eb4c 100644 --- a/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py +++ b/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py @@ -2,15 +2,309 @@ from __future__ import annotations +import fcntl +import os +import re import threading import time from collections import deque from dataclasses import dataclass +from datetime import datetime from enum import Enum +from pathlib import Path +from typing import Any, Mapping import numpy as np +CORE_LAYOUT = { + "observations/qpos": (np.float32, (8,), (256, 8)), + "action": (np.float32, (8,), (256, 8)), + "observations/images/cam_high": ( + np.uint8, + (480, 640, 3), + (1, 480, 640, 3), + ), + "observations/images/cam_right_wrist": ( + np.uint8, + (480, 640, 3), + (1, 480, 640, 3), + ), +} + +DEBUG_LAYOUT = { + "debug/timestamps/control_monotonic_ns": (np.int64, (), (256,)), + "debug/timestamps/feedback_monotonic_ns": (np.int64, (), (256,)), + "debug/timestamps/action_monotonic_ns": (np.int64, (), (256,)), + "debug/timestamps/cam_high_host_monotonic_ns": (np.int64, (), (256,)), + "debug/timestamps/cam_wrist_host_monotonic_ns": (np.int64, (), (256,)), + "debug/timestamps/cam_high_hardware_ms": (np.float64, (), (256,)), + "debug/timestamps/cam_wrist_hardware_ms": (np.float64, (), (256,)), + "debug/timestamps/cam_high_age_ms": (np.float32, (), (256,)), + "debug/timestamps/cam_wrist_age_ms": (np.float32, (), (256,)), + "debug/timestamps/inter_camera_skew_ms": (np.float32, (), (256,)), + "debug/cameras/cam_high_frame_number": (np.uint64, (), (256,)), + "debug/cameras/cam_wrist_frame_number": (np.uint64, (), (256,)), + "debug/control/control_seq": (np.uint64, (), (256,)), + "debug/control/teleop_active": (np.uint8, (), (256,)), + "debug/control/action_valid": (np.uint8, (), (256,)), + "debug/control/command_sent": (np.uint8, (), (256,)), + "debug/control/target_clamped": (np.uint8, (), (256,)), + "debug/control/control_fault": (np.uint8, (), (256,)), + "debug/qp/raw_target": (np.float32, (7,), (256, 7)), + "debug/qp/attempted": (np.uint8, (), (256,)), + "debug/qp/success": (np.uint8, (), (256,)), + "debug/qp/duration_ms": (np.float32, (), (256,)), + "debug/tcp/current_pose": (np.float32, (7,), (256, 7)), + "debug/tcp/raw_target_pose": (np.float32, (7,), (256, 7)), + "debug/tcp/final_target_pose": (np.float32, (7,), (256, 7)), + "debug/tcp/command_velocity": (np.float32, (6,), (256, 6)), + "debug/pico/right_pose": (np.float32, (7,), (256, 7)), + "debug/pico/right_inputs": (np.float32, (6,), (256, 6)), + "debug/pico/left_secondary": (np.uint8, (), (256,)), + "debug/gripper/target_open": (np.uint8, (), (256,)), + "debug/gripper/state_open": (np.uint8, (), (256,)), + "debug/gripper/command_pending": (np.uint8, (), (256,)), + "debug/gripper/command_failed": (np.uint8, (), (256,)), +} + +DATA_LAYOUT = {**CORE_LAYOUT, **DEBUG_LAYOUT} +EPISODE_PATTERN = re.compile(r"^episode_(\d+)\.hdf5$") +PARTIAL_PATTERN = re.compile(r"^episode_(\d+)\.partial\.hdf5$") + + +def _h5py(): + try: + import h5py + except ImportError as exc: + raise RuntimeError("ACT recording requires h5py") from exc + return h5py + + +@dataclass(frozen=True) +class EpisodeMetadata: + joint_names: tuple[str, ...] + joint_lower_limits: np.ndarray + joint_upper_limits: np.ndarray + + def validate(self) -> None: + lower = np.asarray(self.joint_lower_limits) + upper = np.asarray(self.joint_upper_limits) + if len(self.joint_names) != 8: + raise ValueError( + "joint_names must contain seven joints and gripper" + ) + if lower.shape != (7,) or upper.shape != (7,): + raise ValueError("joint limits must have shape (7,)") + if not np.all(np.isfinite(lower)) or not np.all(np.isfinite(upper)): + raise ValueError("joint limits must be finite") + if np.any(lower >= upper): + raise ValueError("joint lower limits must be below upper limits") + + +class EpisodeStore: + def __init__(self, path: Path, root: Any) -> None: + self.path = path + self._root = root + self.count = 0 + + @classmethod + def create(cls, path: Path, metadata: EpisodeMetadata) -> "EpisodeStore": + metadata.validate() + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + root = _h5py().File(path, "x") + root.attrs.update( + { + "sim": np.bool_(False), + "task_name": "tomato_pick", + "sample_rate_hz": 30, + "action_alignment": "same_step_causal", + "arm": "right_rm75", + "episode_status": "recording", + "camera_high_serial": "234222303366", + "camera_right_wrist_serial": "412622272532", + "joint_names": metadata.joint_names, + "joint_lower_limits": np.asarray( + metadata.joint_lower_limits, + dtype=np.float64, + ), + "joint_upper_limits": np.asarray( + metadata.joint_upper_limits, + dtype=np.float64, + ), + "pose_order": "x,y,z,qx,qy,qz,qw", + "right_input_order": ( + "grip,trigger,primary,secondary,axis_x,axis_y" + ), + "interrupted": np.bool_(False), + } + ) + for dataset_path, (dtype, sample_shape, chunks) in DATA_LAYOUT.items(): + parent, separator, name = dataset_path.rpartition("/") + group = root.require_group(parent) if separator else root + group.create_dataset( + name, + shape=(0, *sample_shape), + maxshape=(None, *sample_shape), + chunks=chunks, + dtype=dtype, + ) + return cls(path, root) + + def append(self, values: Mapping[str, Any]) -> None: + if set(values) != set(DATA_LAYOUT): + missing = sorted(set(DATA_LAYOUT) - set(values)) + extra = sorted(set(values) - set(DATA_LAYOUT)) + raise ValueError( + f"sample keys mismatch: missing={missing}, extra={extra}" + ) + + converted = {} + for path, (dtype, shape, _chunks) in DATA_LAYOUT.items(): + value = np.asarray(values[path]) + if value.shape != shape: + raise ValueError(f"{path} must have shape {shape}") + if dtype is np.uint8 and path.startswith("observations/images/"): + if value.dtype != np.uint8: + raise ValueError(f"{path} must have dtype uint8") + elif np.issubdtype(value.dtype, np.floating) and not np.all( + np.isfinite(value) + ): + raise ValueError(f"{path} contains non-finite values") + converted[path] = value.astype(dtype, copy=False) + + new_count = self.count + 1 + for path in DATA_LAYOUT: + self._root[path].resize(new_count, axis=0) + for path, value in converted.items(): + self._root[path][self.count] = value + self.count = new_count + + def truncate(self, count: int) -> None: + if count < 0 or count > self.count: + raise ValueError("truncate count is outside stored samples") + for path in DATA_LAYOUT: + self._root[path].resize(count, axis=0) + self.count = count + + def set_status( + self, + status: str, + *, + reject_reason: str | None = None, + interrupted: bool = False, + ) -> None: + self._root.attrs["episode_status"] = status + self._root.attrs["interrupted"] = np.bool_(interrupted) + if reject_reason is not None: + self._root.attrs["reject_reason"] = reject_reason + + def flush(self) -> None: + self._root.flush() + + def close(self) -> None: + if self._root is not None: + self._root.close() + self._root = None + + +class TaskDirectoryLock: + def __init__(self, directory: Path) -> None: + self.directory = Path(directory) + self._file = None + + def acquire(self) -> None: + if self._file is not None: + raise RuntimeError("task directory lock already acquired") + self.directory.mkdir(parents=True, exist_ok=True) + lock_file = open(self.directory / ".act_recorder.lock", "a+") + try: + fcntl.flock(lock_file, fcntl.LOCK_EX | fcntl.LOCK_NB) + except Exception: + lock_file.close() + raise + self._file = lock_file + + def release(self) -> None: + if self._file is not None: + fcntl.flock(self._file, fcntl.LOCK_UN) + self._file.close() + self._file = None + + +def next_episode_index(directory: Path) -> int: + directory = Path(directory) + if not directory.exists(): + return 0 + indices = [ + int(match.group(1)) + for path in directory.iterdir() + if path.is_file() and (match := EPISODE_PATTERN.fullmatch(path.name)) + ] + return max(indices, default=-1) + 1 + + +def publish_without_overwrite(partial: Path, destination: Path) -> None: + os.link(partial, destination) + Path(partial).unlink() + + +def discard_partial(path: Path) -> None: + path = Path(path) + if not path.name.endswith(".partial.hdf5"): + raise ValueError("only partial episode files can be discarded") + path.unlink(missing_ok=True) + + +def _unique_path(path: Path) -> Path: + candidate = path + suffix = 1 + while candidate.exists(): + candidate = path.with_name(f"{path.stem}_{suffix}{path.suffix}") + suffix += 1 + return candidate + + +def recover_partial_files( + directory: Path, + *, + timestamp: str | None = None, +) -> list[Path]: + directory = Path(directory) + rejected = directory / "rejected" + rejected.mkdir(parents=True, exist_ok=True) + timestamp = timestamp or datetime.now().strftime("%Y%m%dT%H%M%S") + recovered = [] + for path in sorted(directory.iterdir()): + match = PARTIAL_PATTERN.fullmatch(path.name) + if not path.is_file() or match is None: + continue + episode_index = match.group(1) + try: + with _h5py().File(path, "r+") as root: + root.attrs["episode_status"] = "rejected" + root.attrs["reject_reason"] = "crash_recovered" + root.attrs["interrupted"] = np.bool_(True) + except OSError: + destination = _unique_path( + directory + / ( + f"episode_{episode_index}_unreadable_" + f"{timestamp}.partial.hdf5" + ) + ) + path.rename(destination) + else: + destination = _unique_path( + rejected + / f"episode_{episode_index}_crash_recovered_{timestamp}.hdf5" + ) + publish_without_overwrite(path, destination) + recovered.append(destination) + return recovered + + class QualityError(RuntimeError): pass