feat: 集成ACT episode采集节点

This commit is contained in:
2026-08-10 18:22:35 +08:00
parent c1ea1a2816
commit 6008e34b5d
3 changed files with 1009 additions and 1 deletions
+3 -1
View File
@@ -55,7 +55,9 @@ setup(
tests_require=["pytest"], tests_require=["pytest"],
entry_points={ entry_points={
"console_scripts": [ "console_scripts": [
"single_arm_velocity_teleop = xr_rm_teleop.single_arm_velocity_teleop:main", "act_episode_recorder = xr_rm_teleop.act_episode_recorder:main",
"single_arm_velocity_teleop = "
"xr_rm_teleop.single_arm_velocity_teleop:main",
], ],
}, },
) )
@@ -1,3 +1,7 @@
import queue
import threading
from types import SimpleNamespace
import numpy as np import numpy as np
import pytest import pytest
@@ -6,13 +10,16 @@ try:
except ImportError: except ImportError:
h5py = None h5py = None
from xr_rm_interfaces.msg import ActControlSample, XrController
from xr_rm_teleop.act_episode_recorder import ( from xr_rm_teleop.act_episode_recorder import (
NO_ACTION, NO_ACTION,
ActEpisodeRecorder,
ButtonTracker, ButtonTracker,
CameraBuffer, CameraBuffer,
CameraFrame, CameraFrame,
EpisodeMetadata, EpisodeMetadata,
EpisodeStore, EpisodeStore,
EpisodeWriter,
QualityError, QualityError,
QualityLimits, QualityLimits,
RecordingSession, RecordingSession,
@@ -598,3 +605,257 @@ def test_validate_episode_reports_qp_failures_without_rejecting(tmp_path):
assert report.metrics["qp_failure_ratio"] == pytest.approx(2.0 / 3.0) assert report.metrics["qp_failure_ratio"] == pytest.approx(2.0 / 3.0)
assert report.metrics["qp_longest_failure_streak"] == 2 assert report.metrics["qp_longest_failure_streak"] == 2
assert report.metrics["target_clamped_count"] == 2 assert report.metrics["target_clamped_count"] == 2
class _StatusPublisher:
def __init__(self):
self.messages = []
def publish(self, message):
self.messages.append(message.data)
class _Logger:
def info(self, *_args, **_kwargs):
pass
def warn(self, *_args, **_kwargs):
pass
def error(self, *_args, **_kwargs):
pass
def _control_message(seq, control_ns, *, grip=True):
message = ActControlSample()
message.control_seq = seq
message.control_monotonic_ns = control_ns
message.feedback_monotonic_ns = control_ns - 1_000_000
message.action_monotonic_ns = control_ns + 1_000_000 if grip else -1
message.q_actual = [0.0] * 7
message.q_qp_raw = [0.0] * 7
message.q_target = [0.0] * 7
message.joint_lower_limits = [-3.0] * 7
message.joint_upper_limits = [3.0] * 7
message.tcp_current.orientation.w = 1.0
message.tcp_raw_target.orientation.w = 1.0
message.tcp_target.orientation.w = 1.0
message.pico_pose.orientation.w = 1.0
message.pico_grip = grip
message.gripper_target_open = True
message.gripper_state_open = True
message.gripper_state_known = True
message.teleop_active = grip
message.feedback_valid = True
message.action_valid = True
message.command_sent = grip
message.qp_attempted = grip
message.qp_success = grip
return message
def _seed_camera(start_ns, first_number):
camera = SimpleNamespace(
buffer=CameraBuffer(maxlen=4),
last_error=None,
)
image = _image(1)
for index in range(151):
camera.buffer.push(
CameraFrame(
image,
first_number + index,
index * (1000.0 / 30.0),
start_ns
- 5_020_000_000
+ round(index * 5_000_000_000 / 150),
)
)
return camera
def _recorder_for_test(tmp_path):
now_ns = 10_000_000_000
recorder = object.__new__(ActEpisodeRecorder)
recorder._task_dir = tmp_path / "tomato_pick"
recorder._task_dir.mkdir(parents=True)
recorder._quality_limits = _quality_limits()
recorder._min_free_space_bytes = 4 * 1024**3
recorder._controller_timeout_ns = 500_000_000
recorder._camera_warmup_ns = 5_000_000_000
recorder._gripper_completion_timeout_ns = 3_000_000_000
recorder._writer_queue_size = 8
recorder._max_camera_age_ms = 50.0
recorder._max_camera_skew_ms = 50.0
recorder._session = RecordingSession(max_samples=100)
recorder._button_tracker = ButtonTracker(hold_ns=1_000_000_000)
recorder._latest_control = _control_message(99, now_ns - 1_000_000)
recorder._latest_control_received_ns = now_ns - 1_000_000
recorder._right_controller_received_ns = now_ns - 1_000_000
recorder._left_controller_received_ns = now_ns - 1_000_000
recorder._left_secondary = False
recorder._high_camera = _seed_camera(now_ns, 850)
recorder._wrist_camera = _seed_camera(now_ns, 1850)
recorder._camera_start_error = None
recorder._writer = None
recorder._store = None
recorder._partial_path = None
recorder._camera_baselines = None
recorder._saving_deadline_ns = None
recorder._status_pub = _StatusPublisher()
recorder._now_ns = lambda: now_ns
recorder._disk_usage = lambda _path: SimpleNamespace(free=5 * 1024**3)
recorder.get_logger = lambda: _Logger()
return recorder
@requires_h5py
def test_preflight_requires_open_gripper_fresh_inputs_and_disk_space(tmp_path):
recorder = _recorder_for_test(tmp_path)
assert recorder._run_preflight() is None
recorder._latest_control.gripper_state_open = False
assert recorder._run_preflight() == "gripper_not_open"
recorder._latest_control.gripper_state_open = True
recorder._right_controller_received_ns = 0
assert recorder._run_preflight() == "right_controller_stale"
recorder._right_controller_received_ns = recorder._now_ns()
recorder._disk_usage = lambda _path: SimpleNamespace(free=1024)
assert recorder._run_preflight() == "insufficient_disk_space"
def _push_recording_frames(recorder, control_ns, frame_number):
recorder._high_camera.buffer.push(
CameraFrame(
_image(2),
frame_number,
float(frame_number),
control_ns - 2_000_000,
)
)
recorder._wrist_camera.buffer.push(
CameraFrame(
_image(3),
frame_number,
float(frame_number),
control_ns - 3_000_000,
)
)
@requires_h5py
def test_end_to_end_fake_episode_saves_and_returns_idle(tmp_path):
recorder = _recorder_for_test(tmp_path)
recorder._handle_right_b(grip=False)
assert recorder.state is RecordingState.ARMED
start_ns = recorder._now_ns()
frame_number = 1001
for offset, seq in enumerate(range(100, 107)):
control_ns = start_ns + offset * 11_111_111
if (seq - 100) % 3 == 0:
_push_recording_frames(recorder, control_ns, frame_number)
frame_number += 1
recorder._on_control_sample(_control_message(seq, control_ns))
recorder._on_control_sample(
_control_message(107, start_ns + 7 * 11_111_111, grip=False)
)
recorder._handle_right_b(grip=False)
final_ns = start_ns + 8 * 11_111_111
_push_recording_frames(recorder, final_ns, frame_number)
recorder._on_control_sample(_control_message(108, final_ns, grip=False))
assert (tmp_path / "tomato_pick" / "episode_0.hdf5").is_file()
assert recorder.state is RecordingState.IDLE
assert "SAVING" in recorder._status_pub.messages
assert recorder._status_pub.messages[-2:] == ["SAVED", "IDLE"]
@requires_h5py
def test_discard_and_interrupt_only_process_current_partial(tmp_path):
recorder = _recorder_for_test(tmp_path)
recorder._handle_right_b(grip=False)
partial = recorder._partial_path
recorder._discard_current()
assert not partial.exists()
assert recorder.state is RecordingState.IDLE
recorder._handle_right_b(grip=False)
recorder.interrupt_recording("interrupted")
rejected = list((recorder._task_dir / "rejected").glob("*.hdf5"))
assert len(rejected) == 1
with h5py.File(rejected[0], "r") as root:
assert root.attrs["reject_reason"] == "interrupted"
assert root.attrs["interrupted"] == np.bool_(True)
@requires_h5py
def test_a_button_and_camera_error_reject_without_robot_commands(tmp_path):
recorder = _recorder_for_test(tmp_path)
recorder._handle_right_b(grip=False)
recorder._button_tracker.on_right(
False,
False,
False,
recorder._now_ns(),
RecordingState.RECORDING,
)
recorder._session.state = RecordingState.RECORDING
controller = XrController()
controller.primary = True
recorder._on_right_controller(controller)
rejected = list((recorder._task_dir / "rejected").glob("*.hdf5"))
assert len(rejected) == 1
with h5py.File(rejected[0], "r") as root:
assert root.attrs["reject_reason"] == (
"initial_pose_command_during_episode"
)
recorder._handle_right_b(grip=False)
recorder._high_camera.last_error = RuntimeError("usb")
recorder._on_control_sample(
_control_message(100, recorder._now_ns())
)
rejected = list((recorder._task_dir / "rejected").glob("*.hdf5"))
assert len(rejected) == 2
camera_rejected = next(
path for path in rejected if "camera_error" in path.name
)
with h5py.File(camera_rejected, "r") as root:
assert root.attrs["reject_reason"] == "camera_error"
def test_episode_writer_reports_queue_backlog_and_write_error():
started = threading.Event()
release = threading.Event()
class BlockingStore:
def append(self, _sample):
started.set()
assert release.wait(timeout=1.0)
writer = EpisodeWriter(BlockingStore(), queue_size=1)
writer.submit({})
assert started.wait(timeout=1.0)
writer.submit({})
with pytest.raises(queue.Full):
writer.submit({})
release.set()
writer.finish()
assert writer.error is None
class FailingStore:
def append(self, _sample):
raise OSError("disk full")
writer = EpisodeWriter(FailingStore(), queue_size=1)
writer.submit({})
writer.finish()
assert isinstance(writer.error, OSError)
@@ -4,7 +4,9 @@ from __future__ import annotations
import fcntl import fcntl
import os import os
import queue
import re import re
import shutil
import threading import threading
import time import time
from collections import deque from collections import deque
@@ -15,6 +17,11 @@ from pathlib import Path
from typing import Any, Mapping from typing import Any, Mapping
import numpy as np import numpy as np
import rclpy
from rclpy.node import Node
from rclpy.qos import qos_profile_sensor_data
from std_msgs.msg import String
from xr_rm_interfaces.msg import ActControlSample, XrController
CORE_LAYOUT = { CORE_LAYOUT = {
@@ -223,12 +230,69 @@ class EpisodeStore:
def flush(self) -> None: def flush(self) -> None:
self._root.flush() self._root.flush()
def set_attributes(self, values: Mapping[str, Any]) -> None:
self._root.attrs.update(values)
def close(self) -> None: def close(self) -> None:
if self._root is not None: if self._root is not None:
self._root.close() self._root.close()
self._root = None self._root = None
class EpisodeWriter:
_STOP = object()
def __init__(self, store: EpisodeStore, *, queue_size: int) -> None:
if queue_size <= 0:
raise ValueError("queue_size must be positive")
self._store = store
self._queue: queue.Queue = queue.Queue(maxsize=queue_size)
self._lock = threading.Lock()
self._error: Exception | None = None
self._accepting = True
self._finished = False
self._thread = threading.Thread(target=self._run, daemon=True)
self._thread.start()
@property
def error(self) -> Exception | None:
with self._lock:
return self._error
def submit(self, sample: Mapping[str, Any]) -> None:
with self._lock:
if not self._accepting:
raise RuntimeError("episode writer is not accepting samples")
self._queue.put_nowait(sample)
def finish(self) -> None:
with self._lock:
if self._finished:
return
self._accepting = False
self._queue.join()
self._queue.put(self._STOP)
self._queue.join()
self._thread.join()
with self._lock:
self._finished = True
def _run(self) -> None:
while True:
item = self._queue.get()
try:
if item is self._STOP:
return
if self.error is None:
self._store.append(item)
except Exception as exc:
with self._lock:
if self._error is None:
self._error = exc
finally:
self._queue.task_done()
class TaskDirectoryLock: class TaskDirectoryLock:
def __init__(self, directory: Path) -> None: def __init__(self, directory: Path) -> None:
self.directory = Path(directory) self.directory = Path(directory)
@@ -908,3 +972,684 @@ class ButtonTracker:
RecordingState.RECORDING, RecordingState.RECORDING,
) )
self._left_y_fired = False self._left_y_fired = False
def _pose_values(pose: Any) -> np.ndarray:
return np.asarray(
[
pose.position.x,
pose.position.y,
pose.position.z,
pose.orientation.x,
pose.orientation.y,
pose.orientation.z,
pose.orientation.w,
],
dtype=np.float32,
)
def _twist_values(twist: Any) -> np.ndarray:
return np.asarray(
[
twist.linear.x,
twist.linear.y,
twist.linear.z,
twist.angular.x,
twist.angular.y,
twist.angular.z,
],
dtype=np.float32,
)
class ActEpisodeRecorder(Node):
def __init__(self) -> None:
super().__init__("act_episode_recorder")
defaults = {
"output_root": "/home/robot/ACT_Data",
"task_name": "tomato_pick",
"control_sample_topic": "/xr_rm/right_rm75/act_control_sample",
"right_controller_topic": "/xr/right_controller",
"left_controller_topic": "/xr/left_controller",
"status_topic": "/act/recording_status",
"cam_high_serial": "234222303366",
"cam_high_model": "D455",
"cam_right_wrist_serial": "412622272532",
"cam_right_wrist_model": "D405",
"image_width": 640,
"image_height": 480,
"camera_fps": 30,
"camera_warmup_sec": 5.0,
"control_rate_hz": 90.0,
"sample_rate_hz": 30.0,
"min_samples": 60,
"max_samples": 1800,
"min_control_hz": 27.0,
"max_control_gap_ms": 100.0,
"min_camera_fps": 27.0,
"max_drop_ratio": 0.01,
"max_feedback_age_ms": 50.0,
"max_camera_age_ms": 50.0,
"max_camera_skew_ms": 50.0,
"min_free_space_gib": 4.0,
"y_hold_sec": 1.0,
"gripper_completion_timeout_sec": 3.0,
"writer_queue_size": 8,
"controller_timeout_sec": 0.5,
}
for name, value in defaults.items():
self.declare_parameter(name, value)
parameters = {
name: self.get_parameter(name).value for name in defaults
}
self._validate_parameters(parameters)
_h5py()
self._quality_limits = QualityLimits(
min_samples=int(parameters["min_samples"]),
max_samples=int(parameters["max_samples"]),
min_control_hz=float(parameters["min_control_hz"]),
max_control_gap_ms=float(parameters["max_control_gap_ms"]),
min_camera_fps=float(parameters["min_camera_fps"]),
max_drop_ratio=float(parameters["max_drop_ratio"]),
max_feedback_age_ms=float(parameters["max_feedback_age_ms"]),
max_camera_age_ms=float(parameters["max_camera_age_ms"]),
max_camera_skew_ms=float(parameters["max_camera_skew_ms"]),
)
self._max_camera_age_ms = self._quality_limits.max_camera_age_ms
self._max_camera_skew_ms = self._quality_limits.max_camera_skew_ms
self._min_free_space_bytes = int(
float(parameters["min_free_space_gib"]) * 1024**3
)
self._camera_warmup_ns = int(
float(parameters["camera_warmup_sec"]) * 1e9
)
self._controller_timeout_ns = int(
float(parameters["controller_timeout_sec"]) * 1e9
)
self._gripper_completion_timeout_ns = int(
float(parameters["gripper_completion_timeout_sec"]) * 1e9
)
self._writer_queue_size = int(parameters["writer_queue_size"])
self._now_ns = time.monotonic_ns
self._disk_usage = shutil.disk_usage
output_root = Path(str(parameters["output_root"]))
self._task_dir = output_root / str(parameters["task_name"])
self._directory_lock = TaskDirectoryLock(self._task_dir)
self._directory_lock.acquire()
recovered = recover_partial_files(self._task_dir)
if recovered:
self.get_logger().warn(
f"已恢复 {len(recovered)} 个 ACT 临时文件"
)
self._session = RecordingSession(
max_samples=self._quality_limits.max_samples
)
self._button_tracker = ButtonTracker(
hold_ns=int(float(parameters["y_hold_sec"]) * 1e9)
)
self._latest_control: ActControlSample | None = None
self._latest_control_received_ns: int | None = None
self._right_controller_received_ns: int | None = None
self._left_controller_received_ns: int | None = None
self._left_secondary = False
self._writer: EpisodeWriter | None = None
self._store: EpisodeStore | None = None
self._partial_path: Path | None = None
self._camera_baselines: tuple[CameraStats, CameraStats] | None = None
self._saving_deadline_ns: int | None = None
self._high_camera = RealSenseCamera(
str(parameters["cam_high_serial"]),
str(parameters["cam_high_model"]),
)
self._wrist_camera = RealSenseCamera(
str(parameters["cam_right_wrist_serial"]),
str(parameters["cam_right_wrist_model"]),
)
self._camera_start_error: str | None = None
try:
self._high_camera.start()
self._wrist_camera.start()
except Exception as exc:
self._camera_start_error = str(exc)
self.get_logger().error(f"ACT相机启动失败:{exc}")
self._status_pub = self.create_publisher(
String,
str(parameters["status_topic"]),
10,
)
self.create_subscription(
ActControlSample,
str(parameters["control_sample_topic"]),
self._on_control_sample,
qos_profile_sensor_data,
)
self.create_subscription(
XrController,
str(parameters["right_controller_topic"]),
self._on_right_controller,
10,
)
self.create_subscription(
XrController,
str(parameters["left_controller_topic"]),
self._on_left_controller,
10,
)
self._publish_state(RecordingState.IDLE)
@staticmethod
def _validate_parameters(parameters: Mapping[str, Any]) -> None:
if parameters["task_name"] != "tomato_pick":
raise ValueError("task_name must be tomato_pick")
if (
int(parameters["image_width"]) != 640
or int(parameters["image_height"]) != 480
or int(parameters["camera_fps"]) != 30
):
raise ValueError(
"ACT image schema requires RGB8 640x480 at 30 FPS"
)
control_rate = float(parameters["control_rate_hz"])
sample_rate = float(parameters["sample_rate_hz"])
if sample_rate <= 0 or control_rate / sample_rate != 3.0:
raise ValueError("control_rate_hz / sample_rate_hz must equal 3")
for name in (
"min_samples",
"max_samples",
"writer_queue_size",
):
if int(parameters[name]) <= 0:
raise ValueError(f"{name} must be positive")
if int(parameters["min_samples"]) >= int(parameters["max_samples"]):
raise ValueError("min_samples must be below max_samples")
@property
def state(self) -> RecordingState:
return self._session.state
def _publish_state(
self,
state: RecordingState,
reason: str = "",
) -> None:
message = String()
message.data = state.value if not reason else f"{state.value}:{reason}"
self._status_pub.publish(message)
def _run_preflight(self) -> str | None:
now_ns = self._now_ns()
message = self._latest_control
if self._camera_start_error is not None:
return "camera_start_error"
if message is None or self._latest_control_received_ns is None:
return "control_sample_missing"
if (
now_ns - self._latest_control_received_ns
> self._controller_timeout_ns
):
return "control_sample_stale"
if not message.feedback_valid:
return "feedback_invalid"
if message.feedback_age_ms > self._quality_limits.max_feedback_age_ms:
return "feedback_too_old"
if message.control_fault:
return "control_fault"
if not message.gripper_state_known:
return "gripper_state_unknown"
if not message.gripper_state_open:
return "gripper_not_open"
if message.gripper_command_pending:
return "gripper_command_pending"
if message.gripper_command_failed:
return "gripper_command_failed"
if (
self._right_controller_received_ns is None
or now_ns - self._right_controller_received_ns
> self._controller_timeout_ns
):
return "right_controller_stale"
if (
self._left_controller_received_ns is None
or now_ns - self._left_controller_received_ns
> self._controller_timeout_ns
):
return "left_controller_stale"
for name, camera in (
("cam_high", self._high_camera),
("cam_right_wrist", self._wrist_camera),
):
if camera.last_error is not None:
return f"{name}_error"
stats = camera.buffer.stats()
if (
stats.first_host_monotonic_ns is None
or stats.last_host_monotonic_ns is None
or stats.last_host_monotonic_ns
- stats.first_host_monotonic_ns
< self._camera_warmup_ns
):
return f"{name}_warming_up"
if stats.fps < self._quality_limits.min_camera_fps:
return f"{name}_fps"
try:
free_bytes = self._disk_usage(self._task_dir).free
except OSError:
return "output_directory_error"
if free_bytes < self._min_free_space_bytes:
return "insufficient_disk_space"
if not os.access(self._task_dir, os.W_OK):
return "output_directory_not_writable"
return None
def _handle_right_b(self, *, grip: bool) -> None:
if grip:
return
if self.state is RecordingState.IDLE:
self._start_episode()
elif self.state is RecordingState.RECORDING:
self._session.request_finish()
def _start_episode(self) -> None:
reason = self._run_preflight()
if reason is not None:
self._publish_state(RecordingState.IDLE, reason)
self.get_logger().warn(f"ACT录制预检失败:{reason}")
return
assert self._latest_control is not None
episode_index = next_episode_index(self._task_dir)
partial = self._task_dir / f"episode_{episode_index}.partial.hdf5"
metadata = EpisodeMetadata(
joint_names=tuple(
f"omnipic_joint_{index}" for index in range(1, 8)
)
+ ("gripper",),
joint_lower_limits=np.asarray(
self._latest_control.joint_lower_limits,
dtype=np.float64,
),
joint_upper_limits=np.asarray(
self._latest_control.joint_upper_limits,
dtype=np.float64,
),
)
try:
store = EpisodeStore.create(partial, metadata)
except Exception as exc:
self.get_logger().error(f"ACT临时文件创建失败:{exc}")
self._publish_state(RecordingState.IDLE, "disk_write_error")
return
self._store = store
self._writer = EpisodeWriter(
store,
queue_size=self._writer_queue_size,
)
self._partial_path = partial
self._camera_baselines = (
self._high_camera.buffer.stats(),
self._wrist_camera.buffer.stats(),
)
self._saving_deadline_ns = None
self._session.arm()
self._publish_state(RecordingState.ARMED)
def _on_right_controller(self, message: XrController) -> None:
now_ns = self._now_ns()
self._right_controller_received_ns = now_ns
events = self._button_tracker.on_right(
bool(message.primary),
bool(message.secondary),
bool(message.grip),
now_ns,
self.state,
)
if events.reject_reason is not None:
self._reject_current(events.reject_reason)
elif events.right_b:
self._handle_right_b(grip=bool(message.grip))
def _on_left_controller(self, message: XrController) -> None:
now_ns = self._now_ns()
self._left_controller_received_ns = now_ns
self._left_secondary = bool(message.secondary)
events = self._button_tracker.on_left(
bool(message.secondary),
now_ns,
self.state,
)
if events.discard:
self._discard_current()
def _on_control_sample(self, message: ActControlSample) -> None:
self._latest_control = message
self._latest_control_received_ns = self._now_ns()
if self.state is RecordingState.SAVING:
self._continue_finalize(message)
return
if self._writer is not None and self._writer.error is not None:
self._reject_current("disk_write_error")
return
previous_state = self.state
decision = self._session.on_control(
int(message.control_seq),
grip=bool(message.pico_grip),
action_valid=bool(message.action_valid),
command_sent=bool(message.command_sent),
)
if decision.reject_reason is not None:
self._reject_current(decision.reject_reason)
return
if (
previous_state is RecordingState.ARMED
and self.state is RecordingState.RECORDING
):
self._publish_state(RecordingState.RECORDING)
if decision.record_sample:
try:
sample = self._build_sample(message)
assert self._writer is not None
self._writer.submit(sample)
except queue.Full:
self._reject_current("writer_backlog")
return
except QualityError as exc:
self._reject_current(str(exc))
return
if decision.finish:
self._begin_finalize(message)
def _build_sample(self, message: ActControlSample) -> dict[str, Any]:
if self._high_camera.last_error or self._wrist_camera.last_error:
raise QualityError("camera_error")
if not message.gripper_state_known:
raise QualityError("gripper_state_unknown")
high, wrist, high_age, wrist_age, skew = select_camera_pair(
self._high_camera.buffer.snapshot(),
self._wrist_camera.buffer.snapshot(),
int(message.control_monotonic_ns),
max_age_ms=self._max_camera_age_ms,
max_skew_ms=self._max_camera_skew_ms,
)
return {
"observations/qpos": np.asarray(
[*message.q_actual, float(message.gripper_state_open)],
dtype=np.float32,
),
"action": np.asarray(
[*message.q_target, float(message.gripper_target_open)],
dtype=np.float32,
),
"observations/images/cam_high": high.image,
"observations/images/cam_right_wrist": wrist.image,
"debug/timestamps/control_monotonic_ns": (
message.control_monotonic_ns
),
"debug/timestamps/feedback_monotonic_ns": (
message.feedback_monotonic_ns
),
"debug/timestamps/action_monotonic_ns": (
message.action_monotonic_ns
),
"debug/timestamps/cam_high_host_monotonic_ns": (
high.host_monotonic_ns
),
"debug/timestamps/cam_wrist_host_monotonic_ns": (
wrist.host_monotonic_ns
),
"debug/timestamps/cam_high_hardware_ms": (
high.hardware_timestamp_ms
),
"debug/timestamps/cam_wrist_hardware_ms": (
wrist.hardware_timestamp_ms
),
"debug/timestamps/cam_high_age_ms": high_age,
"debug/timestamps/cam_wrist_age_ms": wrist_age,
"debug/timestamps/inter_camera_skew_ms": skew,
"debug/cameras/cam_high_frame_number": high.frame_number,
"debug/cameras/cam_wrist_frame_number": wrist.frame_number,
"debug/control/control_seq": message.control_seq,
"debug/control/teleop_active": message.teleop_active,
"debug/control/action_valid": message.action_valid,
"debug/control/command_sent": message.command_sent,
"debug/control/target_clamped": message.target_clamped,
"debug/control/control_fault": message.control_fault,
"debug/qp/raw_target": np.asarray(
message.q_qp_raw,
dtype=np.float32,
),
"debug/qp/attempted": message.qp_attempted,
"debug/qp/success": message.qp_success,
"debug/qp/duration_ms": message.qp_duration_ms,
"debug/tcp/current_pose": _pose_values(message.tcp_current),
"debug/tcp/raw_target_pose": _pose_values(
message.tcp_raw_target
),
"debug/tcp/final_target_pose": _pose_values(message.tcp_target),
"debug/tcp/command_velocity": _twist_values(
message.tcp_command_velocity
),
"debug/pico/right_pose": _pose_values(message.pico_pose),
"debug/pico/right_inputs": np.asarray(
[
message.pico_grip,
message.pico_trigger,
message.pico_primary,
message.pico_secondary,
*message.pico_axis,
],
dtype=np.float32,
),
"debug/pico/left_secondary": self._left_secondary,
"debug/gripper/target_open": message.gripper_target_open,
"debug/gripper/state_open": message.gripper_state_open,
"debug/gripper/command_pending": (
message.gripper_command_pending
),
"debug/gripper/command_failed": message.gripper_command_failed,
}
@staticmethod
def _interval_camera_metrics(
baseline: CameraStats,
current: CameraStats,
) -> tuple[float, float]:
frames = max(0, current.frame_count - baseline.frame_count)
dropped = max(0, current.dropped_frames - baseline.dropped_frames)
if (
baseline.last_host_monotonic_ns is None
or current.last_host_monotonic_ns is None
):
return 0.0, 1.0
elapsed_ns = (
current.last_host_monotonic_ns
- baseline.last_host_monotonic_ns
)
fps = frames * 1e9 / elapsed_ns if elapsed_ns > 0 else 0.0
expected = frames + dropped
drop_ratio = dropped / expected if expected else 1.0
return fps, drop_ratio
def _write_camera_metrics(self) -> None:
assert self._store is not None
assert self._camera_baselines is not None
high = self._interval_camera_metrics(
self._camera_baselines[0],
self._high_camera.buffer.stats(),
)
wrist = self._interval_camera_metrics(
self._camera_baselines[1],
self._wrist_camera.buffer.stats(),
)
self._store.set_attributes(
{
"camera_high_fps": high[0],
"camera_high_drop_ratio": high[1],
"camera_right_wrist_fps": wrist[0],
"camera_right_wrist_drop_ratio": wrist[1],
}
)
def _begin_finalize(self, message: ActControlSample) -> None:
assert self._writer is not None
assert self._store is not None
self._publish_state(RecordingState.SAVING)
self._writer.finish()
if self._writer.error is not None:
self._reject_current("disk_write_error")
return
crop_count = self._session.candidate_end_count
if crop_count is None:
self._reject_current("missing_release_crop")
return
self._store.truncate(crop_count)
self._write_camera_metrics()
if message.gripper_command_failed:
self._reject_current("gripper_command_failed")
return
if message.gripper_command_pending:
self._saving_deadline_ns = (
self._now_ns() + self._gripper_completion_timeout_ns
)
return
self._complete_save()
def _continue_finalize(self, message: ActControlSample) -> None:
if message.gripper_command_failed:
self._reject_current("gripper_command_failed")
elif not message.gripper_command_pending:
self._complete_save()
elif (
self._saving_deadline_ns is not None
and self._now_ns() >= self._saving_deadline_ns
):
self._reject_current("gripper_command_timeout")
def _complete_save(self) -> None:
assert self._store is not None
assert self._partial_path is not None
partial = self._partial_path
self._store.close()
report = validate_episode(partial, self._quality_limits)
if not report.accepted:
self._reject_closed_partial(
report.reason or "quality_check_failed"
)
return
with _h5py().File(partial, "r+") as root:
root.attrs["episode_status"] = "saved"
root.attrs["interrupted"] = np.bool_(False)
for name, value in report.metrics.items():
root.attrs[name] = value
episode_index = next_episode_index(self._task_dir)
destination = self._task_dir / f"episode_{episode_index}.hdf5"
try:
publish_without_overwrite(partial, destination)
except FileExistsError:
self._reject_closed_partial("episode_number_conflict")
return
self._finish_result(RecordingState.SAVED)
def _reject_closed_partial(
self,
reason: str,
*,
interrupted: bool = False,
) -> None:
assert self._partial_path is not None
with _h5py().File(self._partial_path, "r+") as root:
root.attrs["episode_status"] = "rejected"
root.attrs["reject_reason"] = reason
root.attrs["interrupted"] = np.bool_(interrupted)
rejected = self._task_dir / "rejected"
rejected.mkdir(exist_ok=True)
safe_reason = re.sub(r"[^a-zA-Z0-9_-]", "_", reason)
timestamp = datetime.now().strftime("%Y%m%dT%H%M%S")
match = PARTIAL_PATTERN.fullmatch(self._partial_path.name)
episode_index = match.group(1) if match else "unknown"
destination = _unique_path(
rejected
/ f"episode_{episode_index}_{safe_reason}_{timestamp}.hdf5"
)
publish_without_overwrite(self._partial_path, destination)
self._finish_result(RecordingState.REJECTED, reason)
def _reject_current(
self,
reason: str,
*,
interrupted: bool = False,
) -> None:
if self._partial_path is None:
return
if self._writer is not None:
self._writer.finish()
if self._writer.error is not None:
reason = "disk_write_error"
if self._store is not None:
if self._camera_baselines is not None:
self._write_camera_metrics()
self._store.set_status(
"rejected",
reject_reason=reason,
interrupted=interrupted,
)
self._store.close()
self._reject_closed_partial(reason, interrupted=interrupted)
def _discard_current(self) -> None:
if self._partial_path is None:
return
if self._writer is not None:
self._writer.finish()
if self._store is not None:
self._store.close()
discard_partial(self._partial_path)
self._finish_result(RecordingState.DISCARDED)
def _finish_result(
self,
result: RecordingState,
reason: str = "",
) -> None:
self._session.state = result
self._publish_state(result, reason)
self._session = RecordingSession(
max_samples=self._quality_limits.max_samples
)
self._writer = None
self._store = None
self._partial_path = None
self._camera_baselines = None
self._saving_deadline_ns = None
self._publish_state(RecordingState.IDLE)
def interrupt_recording(self, reason: str = "interrupted") -> None:
if self.state is not RecordingState.IDLE:
self._reject_current(reason, interrupted=True)
def close(self) -> None:
self._high_camera.stop()
self._wrist_camera.stop()
self._directory_lock.release()
def main(args: list[str] | None = None) -> None:
rclpy.init(args=args)
node: ActEpisodeRecorder | None = None
try:
node = ActEpisodeRecorder()
rclpy.spin(node)
except KeyboardInterrupt:
if node is not None:
node.interrupt_recording()
finally:
if node is not None:
node.close()
node.destroy_node()
rclpy.shutdown()