feat: 添加ACT录制状态机

This commit is contained in:
2026-08-10 18:03:25 +08:00
parent e60d620dfe
commit deeee076d7
2 changed files with 393 additions and 0 deletions
@@ -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"
@@ -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