From deeee076d788fc0b969fb4ad9e684484e5663e00 Mon Sep 17 00:00:00 2001 From: YikaiFu-cart Date: Mon, 10 Aug 2026 18:03:25 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0ACT=E5=BD=95=E5=88=B6?= =?UTF-8?q?=E7=8A=B6=E6=80=81=E6=9C=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../test/test_act_episode_recorder.py | 186 ++++++++++++++++ .../xr_rm_teleop/act_episode_recorder.py | 207 ++++++++++++++++++ 2 files changed, 393 insertions(+) create mode 100644 xr_rm_teleop/test/test_act_episode_recorder.py create mode 100644 xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py diff --git a/xr_rm_teleop/test/test_act_episode_recorder.py b/xr_rm_teleop/test/test_act_episode_recorder.py new file mode 100644 index 0000000..fda5988 --- /dev/null +++ b/xr_rm_teleop/test/test_act_episode_recorder.py @@ -0,0 +1,186 @@ +from xr_rm_teleop.act_episode_recorder import ( + NO_ACTION, + ButtonTracker, + RecordingSession, + RecordingState, +) + + +def _recording_session(*, origin_seq=10, max_samples=1800): + session = RecordingSession(max_samples=max_samples) + session.arm() + decision = session.on_control( + origin_seq, + grip=True, + action_valid=True, + command_sent=True, + ) + assert decision.record_sample + return session + + +def test_recording_starts_on_first_sent_grip_action_and_samples_every_third_cycle(): + session = RecordingSession(max_samples=1800) + session.arm() + + assert session.on_control( + 100, + grip=False, + action_valid=True, + command_sent=False, + ) == NO_ACTION + assert session.on_control( + 101, + grip=True, + action_valid=True, + command_sent=False, + ) == NO_ACTION + first = session.on_control( + 102, + grip=True, + action_valid=True, + command_sent=True, + ) + second = session.on_control( + 103, + grip=True, + action_valid=True, + command_sent=True, + ) + session.on_control( + 104, + grip=True, + action_valid=True, + command_sent=True, + ) + third = session.on_control( + 105, + grip=True, + action_valid=True, + command_sent=True, + ) + + assert first.record_sample + assert not second.record_sample + assert third.record_sample + assert session.sample_origin_seq == 102 + + +def test_final_grip_release_marks_crop_point_but_mid_pause_is_kept(): + session = _recording_session(origin_seq=10) + session.on_control(11, grip=True, action_valid=True, command_sent=True) + session.on_control(12, grip=True, action_valid=True, command_sent=True) + session.on_control(13, grip=False, action_valid=True, command_sent=False) + first_crop = session.candidate_end_count + + assert first_crop == 2 + + session.on_control(14, grip=True, action_valid=True, command_sent=False) + assert session.candidate_end_count is None + session.on_control(15, grip=True, action_valid=True, command_sent=True) + session.on_control(16, grip=False, action_valid=True, command_sent=False) + + assert session.candidate_end_count == 3 + + +def test_finish_waits_for_control_sample_after_b_request(): + session = _recording_session(origin_seq=10) + session.on_control(11, grip=False, action_valid=True, command_sent=False) + session.request_finish() + + decision = session.on_control( + 12, + grip=False, + action_valid=True, + command_sent=False, + ) + + assert decision.finish + assert session.candidate_end_count == 1 + + +def test_missing_control_sequence_rejects_recording(): + session = _recording_session(origin_seq=10) + session.on_control(11, grip=True, action_valid=True, command_sent=True) + + decision = session.on_control( + 13, + grip=True, + action_valid=True, + command_sent=True, + ) + + assert decision.reject_reason == "control_sequence_gap" + + +def test_max_samples_rejects_without_stopping_robot(): + session = _recording_session(origin_seq=10, max_samples=2) + session.on_control(11, grip=True, action_valid=True, command_sent=True) + session.on_control(12, grip=True, action_valid=True, command_sent=True) + + decision = session.on_control( + 13, + grip=True, + action_valid=True, + command_sent=True, + ) + + assert decision.reject_reason == "max_duration" + + +def test_right_b_is_ignored_while_grip_is_pressed(): + tracker = ButtonTracker(hold_ns=1_000_000_000) + tracker.on_right(False, False, False, 0, RecordingState.IDLE) + + events = tracker.on_right( + False, + True, + True, + 1, + RecordingState.IDLE, + ) + + assert not events.right_b + + +def test_left_y_requires_new_press_inside_active_recording_state(): + tracker = ButtonTracker(hold_ns=1_000_000_000) + tracker.on_left(True, 0, RecordingState.IDLE) + assert not tracker.on_left( + True, + 2_000_000_000, + RecordingState.RECORDING, + ).discard + tracker.on_left(False, 2_100_000_000, RecordingState.RECORDING) + tracker.on_left(True, 3_000_000_000, RecordingState.RECORDING) + + assert not tracker.on_left( + True, + 3_999_999_999, + RecordingState.RECORDING, + ).discard + assert tracker.on_left( + True, + 4_000_000_000, + RecordingState.RECORDING, + ).discard + assert not tracker.on_left( + True, + 5_000_000_000, + RecordingState.RECORDING, + ).discard + + +def test_right_a_reports_recording_rejection_event(): + tracker = ButtonTracker(hold_ns=1_000_000_000) + tracker.on_right(False, False, False, 0, RecordingState.RECORDING) + + events = tracker.on_right( + True, + False, + False, + 1, + RecordingState.RECORDING, + ) + + assert events.reject_reason == "initial_pose_command_during_episode" diff --git a/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py b/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py new file mode 100644 index 0000000..288060f --- /dev/null +++ b/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py @@ -0,0 +1,207 @@ +"""ALOHA/ACT 风格的右臂 episode 采集节点。""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import Enum + + +class RecordingState(str, Enum): + IDLE = "IDLE" + ARMED = "ARMED" + RECORDING = "RECORDING" + SAVING = "SAVING" + SAVED = "SAVED" + DISCARDED = "DISCARDED" + REJECTED = "REJECTED" + + +@dataclass(frozen=True) +class ControlDecision: + record_sample: bool = False + finish: bool = False + reject_reason: str | None = None + + +NO_ACTION = ControlDecision() + + +class RecordingSession: + def __init__(self, *, max_samples: int) -> None: + if max_samples <= 0: + raise ValueError("max_samples must be positive") + self.max_samples = max_samples + self.state = RecordingState.IDLE + self.sample_origin_seq: int | None = None + self.last_control_seq: int | None = None + self.selected_count = 0 + self.candidate_end_count: int | None = None + self.finish_requested = False + self._last_grip = False + + def arm(self) -> None: + if self.state is not RecordingState.IDLE: + raise RuntimeError("recording session can only arm from IDLE") + self.state = RecordingState.ARMED + self.sample_origin_seq = None + self.last_control_seq = None + self.selected_count = 0 + self.candidate_end_count = None + self.finish_requested = False + self._last_grip = False + + def request_finish(self) -> None: + if self.state is RecordingState.RECORDING: + self.finish_requested = True + + def on_control( + self, + control_seq: int, + *, + grip: bool, + action_valid: bool, + command_sent: bool, + ) -> ControlDecision: + if self.state is RecordingState.ARMED: + if not (grip and action_valid and command_sent): + return NO_ACTION + self.state = RecordingState.RECORDING + self.sample_origin_seq = control_seq + self.last_control_seq = control_seq + self.selected_count = 1 + self._last_grip = True + return ControlDecision(record_sample=True) + + if self.state is not RecordingState.RECORDING: + return NO_ACTION + assert self.sample_origin_seq is not None + assert self.last_control_seq is not None + + if control_seq != self.last_control_seq + 1: + self.state = RecordingState.REJECTED + return ControlDecision(reject_reason="control_sequence_gap") + self.last_control_seq = control_seq + if not action_valid: + self.state = RecordingState.REJECTED + return ControlDecision(reject_reason="invalid_action") + + record_sample = ( + (control_seq - self.sample_origin_seq) % 3 == 0 + ) + if record_sample: + self.selected_count += 1 + + if self._last_grip and not grip: + self.candidate_end_count = self.selected_count + elif not self._last_grip and grip: + self.candidate_end_count = None + self._last_grip = grip + + if ( + self.finish_requested + and not grip + and self.candidate_end_count is not None + ): + self.state = RecordingState.SAVING + return ControlDecision( + record_sample=record_sample, + finish=True, + ) + if self.selected_count >= self.max_samples: + self.state = RecordingState.REJECTED + return ControlDecision( + record_sample=record_sample, + reject_reason="max_duration", + ) + return ControlDecision(record_sample=record_sample) + + +@dataclass(frozen=True) +class ButtonEvents: + right_b: bool = False + right_a: bool = False + discard: bool = False + reject_reason: str | None = None + + +class ButtonTracker: + def __init__(self, *, hold_ns: int) -> None: + if hold_ns <= 0: + raise ValueError("hold_ns must be positive") + self.hold_ns = hold_ns + self._right_a: bool | None = None + self._right_b: bool | None = None + self._left_y: bool | None = None + self._left_y_started_ns: int | None = None + self._left_y_eligible = False + self._left_y_fired = False + + def on_right( + self, + primary: bool, + secondary: bool, + grip: bool, + now_ns: int, + state: RecordingState, + ) -> ButtonEvents: + del now_ns + if self._right_a is None or self._right_b is None: + self._right_a = primary + self._right_b = secondary + return ButtonEvents() + + a_rising = primary and not self._right_a + b_rising = secondary and not self._right_b + self._right_a = primary + self._right_b = secondary + reject_reason = None + if a_rising and state is RecordingState.RECORDING: + reject_reason = "initial_pose_command_during_episode" + return ButtonEvents( + right_b=b_rising and not grip, + right_a=a_rising, + reject_reason=reject_reason, + ) + + def on_left( + self, + secondary: bool, + now_ns: int, + state: RecordingState, + ) -> ButtonEvents: + if self._left_y is None: + self._left_y = secondary + if secondary: + self._start_left_y(now_ns, state) + return ButtonEvents() + + if secondary and not self._left_y: + self._start_left_y(now_ns, state) + elif not secondary and self._left_y: + self._left_y_started_ns = None + self._left_y_eligible = False + self._left_y_fired = False + self._left_y = secondary + + if ( + secondary + and self._left_y_eligible + and not self._left_y_fired + and self._left_y_started_ns is not None + and now_ns - self._left_y_started_ns >= self.hold_ns + ): + self._left_y_fired = True + return ButtonEvents(discard=True) + return ButtonEvents() + + def _start_left_y( + self, + now_ns: int, + state: RecordingState, + ) -> None: + self._left_y_started_ns = now_ns + self._left_y_eligible = state in ( + RecordingState.ARMED, + RecordingState.RECORDING, + ) + self._left_y_fired = False