From b5d46c4edd8d36dd555f6b82782558a54a88181c Mon Sep 17 00:00:00 2001 From: Saberlve Date: Fri, 7 Aug 2026 17:49:10 +0800 Subject: [PATCH] Add data replay --- .gitignore | 2 +- README_ZH.md | 16 + .../xarm7_manual_record_config.yaml | 4 +- .../manual_mode/xarm_manual_mode_config.yaml | 14 - pyproject.toml | 2 + .../robots/uf_robot/uf_robot.py | 76 ++++- .../robots/uf_robot/uf_robot_config.py | 3 + .../scripts/uf_lerobot_replay.py | 272 +++++++++++++++++ tests/test_lerobot_replay.py | 273 ++++++++++++++++++ uv.lock | 2 + 10 files changed, 635 insertions(+), 29 deletions(-) delete mode 100644 config/manual_mode/xarm_manual_mode_config.yaml create mode 100644 src/lerobot_robot_ufactory/scripts/uf_lerobot_replay.py create mode 100644 tests/test_lerobot_replay.py diff --git a/.gitignore b/.gitignore index 32dee99..8aa32ef 100644 --- a/.gitignore +++ b/.gitignore @@ -89,4 +89,4 @@ models/ *.xvcd ufactory_usage/ .history/ -xarm7_manual_datas \ No newline at end of file +datasets/ \ No newline at end of file diff --git a/README_ZH.md b/README_ZH.md index 9e1d02c..c4d0bff 100644 --- a/README_ZH.md +++ b/README_ZH.md @@ -178,6 +178,22 @@ uv run uf-lerobot-record --config_path path/to/config.yaml --resume true # uv run uf-lerobot-record --config_path config/umi/xarm6_umi_record_config.yaml ``` +### 4. 数据重放 + +`datasets/xarm7_manual_replay` 是人工拖拽录制的 LeRobot 数据集。回放脚本使用其中的 +`observation.state`,将 7 个关节弧度值和归一化夹爪位置作为**绝对目标值**发送给 xArm7, +不会将相邻帧相减,也不会累加成相对动作。脚本默认按数据集的 30 FPS 播放一个 episode。 + +启动前会要求确认,连接后会先自动移动到 xArm SDK 初始点;播放结束后保持最后一帧姿态并断开连接: + +```bash +uv run uf-lerobot-replay \ + --dataset-root /home/wsx/code/lerobot_robot_ufactory/datasets/xarm7_manual_replay \ + --robot-ip 192.168.1.245 +``` + +无人值守运行时可以使用 `--yes` 跳过确认。执行前请确认机械臂工作空间无障碍物,且数据中的初始姿态与当前设备匹配。 + ### 5. Lerobot训练 采集数据后,使用 LeRobot 训练管道进行模仿学习训练。 diff --git a/config/manual_mode/xarm7_manual_record_config.yaml b/config/manual_mode/xarm7_manual_record_config.yaml index 6f67570..acab5fd 100644 --- a/config/manual_mode/xarm7_manual_record_config.yaml +++ b/config/manual_mode/xarm7_manual_record_config.yaml @@ -24,7 +24,7 @@ robot: fps: 30 dataset: - root: "/home/wsx/code/lerobot_robot_ufactory/xarm7_manual_datas" + root: "/home/wsx/code/lerobot_robot_ufactory/datasets/xarm7_manual_replay_pick_pen" repo_id: "ufactory/xarm7_manual_datas" # Task description stored with each recorded frame. single_task: "Describe the task being demonstrated." @@ -37,3 +37,5 @@ dataset: # Store camera observations as videos. video: true push_to_hub: false + + diff --git a/config/manual_mode/xarm_manual_mode_config.yaml b/config/manual_mode/xarm_manual_mode_config.yaml deleted file mode 100644 index 8a18252..0000000 --- a/config/manual_mode/xarm_manual_mode_config.yaml +++ /dev/null @@ -1,14 +0,0 @@ -# xArm manual/joint teaching mode test configuration. -robot_ip: "192.168.1.245" - -# true: enter mode 2 and wait for Enter; false: restore mode 0. -manual_mode: true - -# xArm teach sensitivity, valid range: 1-5. -teach_sensitivity: 3 - -# Read the initial point saved in xArm Studio and return to it after pressing Enter. -return_to_initial: true - -# Return speed in degrees per second. -reset_speed: 30 diff --git a/pyproject.toml b/pyproject.toml index 6a3a8b0..80c895f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,6 +25,7 @@ classifiers = [ ] dependencies = [ "numpy>=1.24", + "pyarrow>=14.0", "pyyaml", "lerobot[intelrealsense]==0.4.3", "xarm-python-sdk", @@ -35,6 +36,7 @@ dependencies = [ uf-robot-teleop = "lerobot_robot_ufactory.scripts.uf_robot_teleop:main" uf-lerobot-record = "lerobot_robot_ufactory.scripts.uf_lerobot_record:main" uf-lerobot-eval = "lerobot_robot_ufactory.scripts.uf_lerobot_eval:main" +uf-lerobot-replay = "lerobot_robot_ufactory.scripts.uf_lerobot_replay:main" uf-vive-calibrate = "lerobot_robot_ufactory.scripts.vive_calibrate:main" uf-camera-view = "lerobot_robot_ufactory.scripts.uf_camera_view:main" uf-camera-test = "lerobot_robot_ufactory.scripts.uf_camera_test:main" diff --git a/src/lerobot_robot_ufactory/robots/uf_robot/uf_robot.py b/src/lerobot_robot_ufactory/robots/uf_robot/uf_robot.py index 287aef5..99265a4 100644 --- a/src/lerobot_robot_ufactory/robots/uf_robot/uf_robot.py +++ b/src/lerobot_robot_ufactory/robots/uf_robot/uf_robot.py @@ -223,8 +223,16 @@ class UFRobot(Robot, Thread): if self._initial_point is None: raise RuntimeError("xArm initial point has not been loaded") - self.real_arm.set_mode(0) - self.real_arm.set_state(0) + # The controller requires motion to be enabled again after an + # emergency stop has been released, before any reset motion command. + code = self.real_arm.motion_enable(enable=True) + self._check_motion_code("motion_enable", code) + code = self.real_arm.clean_error() + self._check_motion_code("clean_error", code) + code = self.real_arm.set_mode(0) + self._check_motion_code("set_mode(0)", code) + code = self.real_arm.set_state(0) + self._check_motion_code("set_state(0)", code) code = self.real_arm.set_servo_angle( angle=self._initial_point, speed=ROBOT_RESET_SPEED_DEG, @@ -271,7 +279,9 @@ class UFRobot(Robot, Thread): return if self._control_space == "joint": - self.real_arm.set_mode(6) + code = self.real_arm.set_mode(self.config.joint_command_mode) + if code != 0: + raise RuntimeError(f"set_mode({self.config.joint_command_mode}) failed, code={code}") elif self._control_space == "cartesian": self.real_arm.set_mode(7) else: @@ -430,6 +440,20 @@ class UFRobot(Robot, Thread): modbus_datas = [0x09, 0x10, 0x03, 0xE8, 0x00, 0x03, 0x06, 0x09, 0x00, 0x00, grippos, self._gripper_param.speed, self._gripper_param.force] self.real_arm.getset_tgpio_modbus_data(modbus_datas) + def _motion_status(self) -> str: + """Return controller state details for a failed motion command.""" + arm = self.real_arm + mode = getattr(arm, "mode", "unknown") + state = getattr(arm, "state", "unknown") + error_code = getattr(arm, "error_code", "unknown") + warn_code = getattr(arm, "warn_code", "unknown") + return f"mode={mode}, state={state}, error_code={error_code}, warn_code={warn_code}" + + def _check_motion_code(self, command: str, code: int) -> None: + """Fail loudly when the SDK rejects a joint command.""" + if code != 0: + raise RuntimeError(f"{command} failed, code={code}, {self._motion_status()}") + def send_action(self, action: dict) -> np.ndarray: if not self._is_connected: raise ConnectionError() @@ -458,17 +482,43 @@ class UFRobot(Robot, Thread): for i in range(self._dof): cmd_list[i] = action[f"{self.prefix}J{i+1}.pos"] - # TODO: make mode 6 compatible with wait=True - if wait_== False and self.real_arm.mode != 6: - self.real_arm.set_mode(6) - self.real_arm.set_state(0) - time.sleep(0.1) - elif wait_ and self.real_arm.mode != 0: - self.real_arm.set_mode(0) - self.real_arm.set_state(0) - time.sleep(0.1) + if self.config.joint_command_mode == 1: + # set_servo_angle_j is an absolute target command. It is the + # SDK's high-frequency interface and executes only the latest + # target, so it must be used with servo motion mode (1). + if self.real_arm.mode != 1: + code = self.real_arm.set_mode(1) + self._check_motion_code("set_mode(1)", code) + code = self.real_arm.set_state(0) + self._check_motion_code("set_state(0)", code) + time.sleep(0.1) + code = self.real_arm.set_servo_angle_j( + cmd_list[:self._dof], speed=jnt_spd, is_radian=True + ) + self._check_motion_code("set_servo_angle_j", code) + else: + # The legacy mode-6 path uses the absolute move_joint API. + # The first blocking command must be sent in position mode. + if wait_ == False and self.real_arm.mode != 6: + code = self.real_arm.set_mode(6) + self._check_motion_code("set_mode(6)", code) + code = self.real_arm.set_state(0) + self._check_motion_code("set_state(0)", code) + time.sleep(0.1) + elif wait_ and self.real_arm.mode != 0: + code = self.real_arm.set_mode(0) + self._check_motion_code("set_mode(0)", code) + code = self.real_arm.set_state(0) + self._check_motion_code("set_state(0)", code) + time.sleep(0.1) - self.real_arm.set_servo_angle(angle=cmd_list[:self._dof], speed=jnt_spd, is_radian=True, wait=wait_) + code = self.real_arm.set_servo_angle( + angle=cmd_list[:self._dof], + speed=jnt_spd, + is_radian=True, + wait=wait_, + ) + self._check_motion_code("set_servo_angle", code) elif self._control_space == "cartesian": # unit: mm? lin_spd = self._max_linear_velocity diff --git a/src/lerobot_robot_ufactory/robots/uf_robot/uf_robot_config.py b/src/lerobot_robot_ufactory/robots/uf_robot/uf_robot_config.py index f5da02e..97b76f3 100644 --- a/src/lerobot_robot_ufactory/robots/uf_robot/uf_robot_config.py +++ b/src/lerobot_robot_ufactory/robots/uf_robot/uf_robot_config.py @@ -20,6 +20,7 @@ class UFRobotConfig(RobotConfig): manual_mode: bool = False # xArm joint teaching mode; records state and optional gripper actions manual_gripper_speed: float = 0.5 # normalized gripper position per second in manual mode teach_sensitivity: int | None = None # xArm teaching sensitivity, valid range: 1-5 + joint_command_mode: int = 6 # 1: servo-angle-j, 6: online trajectory planning # start_joints and start_tcp_pose are intentionally disabled. # Reset uses the xArm SDK initial_point instead of configuration poses. max_joint_velocity: int = 90 # °/s, only effective in joint control mode @@ -36,3 +37,5 @@ class UFRobotConfig(RobotConfig): raise ValueError("teach_sensitivity must be between 1 and 5") if self.manual_gripper_speed < 0: raise ValueError("manual_gripper_speed must be non-negative") + if self.control_space == "joint" and self.joint_command_mode not in (1, 6): + raise ValueError("joint_command_mode must be 1 or 6 for joint control") diff --git a/src/lerobot_robot_ufactory/scripts/uf_lerobot_replay.py b/src/lerobot_robot_ufactory/scripts/uf_lerobot_replay.py new file mode 100644 index 0000000..7a7a2af --- /dev/null +++ b/src/lerobot_robot_ufactory/scripts/uf_lerobot_replay.py @@ -0,0 +1,272 @@ +"""Replay absolute joint states from a LeRobot dataset on an xArm robot.""" + +from __future__ import annotations + +import argparse +import json +import math +import sys +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Sequence + +import pyarrow.parquet as parquet + +import lerobot_robot_ufactory # noqa: F401 +from lerobot.utils.robot_utils import precise_sleep +from lerobot_robot_ufactory.robots.uf_robot.uf_robot import UFRobot +from lerobot_robot_ufactory.robots.uf_robot.uf_robot_config import UFRobotConfig + + +JOINT_STATE_NAMES = tuple(f"J{index}.pos" for index in range(1, 8)) +STATE_NAMES = JOINT_STATE_NAMES + ("gripper.pos",) +STATE_FEATURE = "observation.state" +DEFAULT_DATASET_ROOT = Path("datasets/xarm7_manual_replay") +DEFAULT_ROBOT_IP = "192.168.1.245" + + +@dataclass(frozen=True) +class ReplayEpisode: + """Validated, ordered absolute states for one dataset episode.""" + + fps: float + episode_index: int + frame_indices: tuple[int, ...] + states: tuple[tuple[float, ...], ...] + + +def _load_info(dataset_root: Path) -> dict[str, Any]: + info_path = dataset_root / "meta" / "info.json" + if not info_path.is_file(): + raise ValueError(f"LeRobot metadata file does not exist: {info_path}") + + try: + with info_path.open("r", encoding="utf-8") as info_file: + info = json.load(info_file) + except (OSError, json.JSONDecodeError) as exc: + raise ValueError(f"Could not read LeRobot metadata: {info_path}") from exc + + if not isinstance(info, dict): + raise ValueError(f"LeRobot metadata must contain a JSON object: {info_path}") + return info + + +def _validate_state_feature(info: dict[str, Any]) -> float: + try: + fps = float(info["fps"]) + feature = info["features"][STATE_FEATURE] + names = tuple(feature["names"]) + shape = tuple(feature["shape"]) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError( + f"Dataset metadata must define {STATE_FEATURE!r} and a positive FPS" + ) from exc + + if not math.isfinite(fps) or fps <= 0: + raise ValueError(f"Dataset FPS must be a positive finite number, got {fps!r}") + if names != STATE_NAMES or shape != (len(STATE_NAMES),): + raise ValueError( + f"{STATE_FEATURE!r} must contain absolute fields {list(STATE_NAMES)!r}; " + f"got names={list(names)!r}, shape={list(shape)!r}" + ) + return fps + + +def _read_episode_rows(dataset_root: Path, episode_index: int) -> list[dict[str, Any]]: + data_files = sorted((dataset_root / "data").glob("chunk-*/file-*.parquet")) + if not data_files: + raise ValueError(f"No LeRobot data files found below {dataset_root / 'data'}") + + rows: list[dict[str, Any]] = [] + columns = [STATE_FEATURE, "episode_index", "frame_index"] + for data_file in data_files: + try: + table = parquet.read_table(data_file, columns=columns) + except Exception as exc: + raise ValueError(f"Could not read LeRobot data file: {data_file}") from exc + rows.extend( + row for row in table.to_pylist() if int(row["episode_index"]) == episode_index + ) + + if not rows: + raise ValueError(f"Episode {episode_index} does not exist in {dataset_root}") + rows.sort(key=lambda row: int(row["frame_index"])) + return rows + + +def load_replay_episode(dataset_root: Path, episode_index: int = 0) -> ReplayEpisode: + """Load and validate one episode as absolute robot target states.""" + + dataset_root = Path(dataset_root).expanduser().resolve() + if not dataset_root.is_dir(): + raise ValueError(f"Dataset directory does not exist: {dataset_root}") + if episode_index < 0: + raise ValueError(f"Episode index must be non-negative, got {episode_index}") + + info = _load_info(dataset_root) + fps = _validate_state_feature(info) + rows = _read_episode_rows(dataset_root, episode_index) + + frame_indices = tuple(int(row["frame_index"]) for row in rows) + expected_indices = tuple(range(len(rows))) + if frame_indices != expected_indices: + raise ValueError( + f"Episode {episode_index} frame_index must be contiguous from 0; " + f"got first={frame_indices[0]}, last={frame_indices[-1]}" + ) + + states: list[tuple[float, ...]] = [] + for row_index, row in enumerate(rows): + raw_state = row[STATE_FEATURE] + if not isinstance(raw_state, (list, tuple)) or len(raw_state) != len(STATE_NAMES): + raise ValueError( + f"Episode {episode_index}, frame {row_index} must contain " + f"{len(STATE_NAMES)} state values" + ) + state = tuple(float(value) for value in raw_state) + if not all(math.isfinite(value) for value in state): + raise ValueError( + f"Episode {episode_index}, frame {row_index} contains a non-finite value" + ) + states.append(state) + + return ReplayEpisode( + fps=fps, + episode_index=episode_index, + frame_indices=frame_indices, + states=tuple(states), + ) + + +def state_to_robot_action(state: Sequence[float]) -> dict[str, float]: + """Map one absolute dataset state to the xArm absolute command format.""" + + if len(state) != len(STATE_NAMES): + raise ValueError(f"Expected {len(STATE_NAMES)} state values, got {len(state)}") + values = tuple(float(value) for value in state) + if not all(math.isfinite(value) for value in values): + raise ValueError("Robot state contains a non-finite value") + # These values are absolute positions. Do not subtract the previous state. + return dict(zip(STATE_NAMES, values, strict=True)) + + +def replay_episode(robot: Any, episode: ReplayEpisode, show_progress: bool = True) -> None: + """Send every absolute state once at the dataset FPS.""" + + period_s = 1.0 / episode.fps + next_deadline = time.perf_counter() + total_frames = len(episode.states) + + for frame_number, state in enumerate(episode.states, start=1): + robot.send_action(state_to_robot_action(state)) + if show_progress: + print(f"\rReplaying frame {frame_number}/{total_frames}", end="", flush=True) + + next_deadline += period_s + precise_sleep(max(next_deadline - time.perf_counter(), 0.0)) + + if show_progress: + print() + + +def _build_robot(robot_ip: str) -> UFRobot: + config = UFRobotConfig( + id="xarm7_replay_robot", + robot_ip=robot_ip, + robot_dof=7, + control_space="joint", + joint_command_mode=1, + gripper_type=1, + manual_mode=False, + cameras={}, + ) + return UFRobot(config) + + +def _confirm_start(episode: ReplayEpisode, robot_ip: str, skip_confirmation: bool) -> bool: + first_state = state_to_robot_action(episode.states[0]) + last_state = state_to_robot_action(episode.states[-1]) + duration_s = (len(episode.states) - 1) / episode.fps + print(f"Robot: xArm7 at {robot_ip}") + print( + f"Episode {episode.episode_index}: {len(episode.states)} frames, " + f"{episode.fps:g} FPS, about {duration_s:.2f} seconds" + ) + print(f"First absolute state: {first_state}") + print(f"Last absolute state: {last_state}") + print("The robot will first move to its xArm SDK initial point.") + print("After replay it will stop at the last state and disconnect.") + + if skip_confirmation: + return True + if not sys.stdin.isatty(): + raise RuntimeError("Interactive confirmation is required; use --yes to continue") + return input("Type 'yes' to connect and start replay: ").strip().lower() == "yes" + + +def _disconnect_quietly(robot: UFRobot) -> None: + if getattr(robot, "real_arm", None) is None: + return + try: + robot.disconnect() + except Exception as exc: + print(f"Warning: failed to disconnect robot cleanly: {exc}", file=sys.stderr) + + +def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Replay absolute observation.state joint positions on an xArm7." + ) + parser.add_argument( + "--dataset-root", + type=Path, + default=DEFAULT_DATASET_ROOT, + help=f"LeRobot dataset root (default: {DEFAULT_DATASET_ROOT})", + ) + parser.add_argument( + "--robot-ip", + default=DEFAULT_ROBOT_IP, + help=f"xArm controller IP (default: {DEFAULT_ROBOT_IP})", + ) + parser.add_argument( + "--episode-index", + type=int, + default=0, + help="Episode to replay (default: 0)", + ) + parser.add_argument( + "--yes", + action="store_true", + help="Skip the interactive confirmation before connecting to the robot", + ) + return parser.parse_args(argv) + + +def main(argv: Sequence[str] | None = None) -> int: + args = parse_args(argv) + try: + episode = load_replay_episode(args.dataset_root, args.episode_index) + if not _confirm_start(episode, args.robot_ip, args.yes): + print("Replay cancelled before connecting to the robot.") + return 0 + + robot = _build_robot(args.robot_ip) + try: + robot.connect() + print("Robot connected and moved to the SDK initial point.") + replay_episode(robot, episode) + print("Replay complete. The robot will remain at the last state.") + finally: + _disconnect_quietly(robot) + except KeyboardInterrupt: + print("\nReplay interrupted by user.", file=sys.stderr) + return 130 + except (RuntimeError, ValueError, OSError) as exc: + print(f"Replay failed: {exc}", file=sys.stderr) + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/test_lerobot_replay.py b/tests/test_lerobot_replay.py new file mode 100644 index 0000000..ad13289 --- /dev/null +++ b/tests/test_lerobot_replay.py @@ -0,0 +1,273 @@ +import json + +import pyarrow as pa +import pyarrow.parquet as parquet +import pytest + +from lerobot_robot_ufactory.scripts import uf_lerobot_replay as replay_module +from lerobot_robot_ufactory.scripts.uf_lerobot_replay import ( + JOINT_STATE_NAMES, + ReplayEpisode, + STATE_NAMES, + load_replay_episode, + replay_episode, + state_to_robot_action, +) + + +def _write_dataset(tmp_path, states, episode_indices=None, frame_indices=None, names=None): + dataset_root = tmp_path / "dataset" + (dataset_root / "data" / "chunk-000").mkdir(parents=True) + (dataset_root / "meta").mkdir() + + episode_indices = episode_indices or [0] * len(states) + frame_indices = frame_indices or list(range(len(states))) + names = names or list(STATE_NAMES) + info = { + "fps": 30, + "features": { + "observation.state": { + "dtype": "float32", + "names": names, + "shape": [len(names)], + } + }, + } + (dataset_root / "meta" / "info.json").write_text( + json.dumps(info), encoding="utf-8" + ) + table = pa.table( + { + "observation.state": states, + "episode_index": episode_indices, + "frame_index": frame_indices, + } + ) + parquet.write_table(table, dataset_root / "data" / "chunk-000" / "file-000.parquet") + return dataset_root + + +def test_load_replay_episode_sorts_and_maps_absolute_states(tmp_path): + state_a = [0.1, -0.2, 0.3, -0.4, 0.5, -0.6, 0.7, 1.00375] + state_b = [0.2, -0.3, 0.4, -0.5, 0.6, -0.7, 0.8, 0.4] + dataset_root = _write_dataset( + tmp_path, + states=[state_b, state_a, [9.0] * 8], + episode_indices=[0, 0, 1], + frame_indices=[1, 0, 0], + ) + + episode = load_replay_episode(dataset_root, episode_index=0) + + assert episode.fps == 30 + assert episode.frame_indices == (0, 1) + assert episode.states == (tuple(state_a), tuple(state_b)) + assert state_to_robot_action(state_a) == dict(zip(STATE_NAMES, state_a, strict=True)) + + +def test_replay_sends_absolute_values_without_delta_accumulation(monkeypatch): + class FakeRobot: + def __init__(self): + self.actions = [] + + def send_action(self, action): + self.actions.append(action.copy()) + + monkeypatch.setattr(replay_module, "precise_sleep", lambda _: None) + state_a = (0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8) + state_b = (0.2, 0.4, 0.6, 0.8, 1.0, 1.2, 1.4, 0.9) + episode = ReplayEpisode( + fps=30, + episode_index=0, + frame_indices=(0, 1), + states=(state_a, state_b), + ) + robot = FakeRobot() + + replay_episode(robot, episode, show_progress=False) + + assert robot.actions == [ + dict(zip(STATE_NAMES, state_a, strict=True)), + dict(zip(STATE_NAMES, state_b, strict=True)), + ] + assert robot.actions[1]["J1.pos"] == state_b[0] + + +def test_ufactory_robot_sends_absolute_radian_joint_targets(monkeypatch): + from lerobot_robot_ufactory.robots.uf_robot import uf_robot as uf_robot_module + from lerobot_robot_ufactory.robots.uf_robot.uf_robot_config import UFRobotConfig + + class FakeArm: + error_code = 0 + mode = 6 + + def __init__(self): + self.calls = [] + + def set_mode(self, mode): + self.calls.append(("set_mode", mode)) + self.mode = mode + return 0 + + def set_state(self, state): + self.calls.append(("set_state", state)) + return 0 + + def set_servo_angle(self, **kwargs): + self.calls.append(("set_servo_angle", kwargs)) + return 0 + + monkeypatch.setattr(uf_robot_module.time, "sleep", lambda _: None) + robot = uf_robot_module.UFRobot( + UFRobotConfig( + id="replay-test", + robot_dof=7, + control_space="joint", + gripper_type=0, + cameras={}, + ) + ) + arm = FakeArm() + robot.real_arm = arm + robot._is_connected = True + target = state_to_robot_action((0.1, -0.2, 0.3, -0.4, 0.5, -0.6, 0.7, 0.8)) + + robot.send_action(target) + + servo_calls = [kwargs for name, kwargs in arm.calls if name == "set_servo_angle"] + assert servo_calls == [ + { + "angle": [0.1, -0.2, 0.3, -0.4, 0.5, -0.6, 0.7], + "speed": 0.2, + "is_radian": True, + "wait": True, + } + ] + + +def test_reset_to_initial_enables_robot_before_motion(monkeypatch): + from lerobot_robot_ufactory.robots.uf_robot import uf_robot as uf_robot_module + from lerobot_robot_ufactory.robots.uf_robot.uf_robot_config import UFRobotConfig + + class FakeArm: + mode = 0 + state = 0 + error_code = 0 + warn_code = 0 + + def __init__(self): + self.calls = [] + + def motion_enable(self, enable=True): + self.calls.append(("motion_enable", enable)) + return 0 + + def clean_error(self): + self.calls.append(("clean_error",)) + return 0 + + def set_mode(self, mode): + self.calls.append(("set_mode", mode)) + self.mode = mode + return 0 + + def set_state(self, state): + self.calls.append(("set_state", state)) + self.state = state + return 0 + + def set_servo_angle(self, **kwargs): + self.calls.append(("set_servo_angle", kwargs)) + return 0 + + monkeypatch.setattr(uf_robot_module.time, "sleep", lambda _: None) + robot = uf_robot_module.UFRobot( + UFRobotConfig( + id="replay-reset-test", + robot_dof=7, + control_space="joint", + gripper_type=0, + cameras={}, + ) + ) + arm = FakeArm() + robot.real_arm = arm + robot._is_connected = True + robot._initial_point = [0.0] * 7 + robot.configure = lambda: None + + robot.reset_to_initial() + + assert [name for name, *_ in arm.calls] == [ + "motion_enable", + "clean_error", + "set_mode", + "set_state", + "set_servo_angle", + ] + assert arm.calls[0] == ("motion_enable", True) + + +def test_ufactory_robot_replay_path_sends_absolute_servo_j_targets(): + from lerobot_robot_ufactory.robots.uf_robot import uf_robot as uf_robot_module + from lerobot_robot_ufactory.robots.uf_robot.uf_robot_config import UFRobotConfig + + class FakeArm: + error_code = 0 + mode = 1 + state = 0 + warn_code = 0 + + def __init__(self): + self.calls = [] + + def set_servo_angle_j(self, angles, **kwargs): + self.calls.append(("set_servo_angle_j", angles, kwargs)) + return 0 + + robot = uf_robot_module.UFRobot( + UFRobotConfig( + id="replay-servoj-test", + robot_dof=7, + control_space="joint", + joint_command_mode=1, + gripper_type=0, + cameras={}, + ) + ) + arm = FakeArm() + robot.real_arm = arm + robot._is_connected = True + + target = state_to_robot_action((0.1, -0.2, 0.3, -0.4, 0.5, -0.6, 0.7, 0.8)) + robot.send_action(target) + + assert arm.calls == [ + ( + "set_servo_angle_j", + [0.1, -0.2, 0.3, -0.4, 0.5, -0.6, 0.7], + {"speed": 0.2, "is_radian": True}, + ) + ] + + +def test_load_replay_episode_rejects_wrong_state_schema(tmp_path): + states = [[0.0] * 8] + dataset_root = _write_dataset(tmp_path, states, names=list(JOINT_STATE_NAMES) + ["wrong"]) + + with pytest.raises(ValueError, match="absolute fields"): + load_replay_episode(dataset_root) + + +def test_load_replay_episode_rejects_non_contiguous_frames(tmp_path): + dataset_root = _write_dataset(tmp_path, [[0.0] * 8, [1.0] * 8], frame_indices=[0, 2]) + + with pytest.raises(ValueError, match="contiguous"): + load_replay_episode(dataset_root) + + +def test_state_to_robot_action_rejects_non_finite_values(): + state = [0.0] * 7 + [float("nan")] + + with pytest.raises(ValueError, match="non-finite"): + state_to_robot_action(state) diff --git a/uv.lock b/uv.lock index b7fdb99..ce8071d 100644 --- a/uv.lock +++ b/uv.lock @@ -1223,6 +1223,7 @@ dependencies = [ { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, { name = "numpy", version = "2.5.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, { name = "opencv-python" }, + { name = "pyarrow" }, { name = "pyyaml" }, { name = "xarm-python-sdk" }, ] @@ -1252,6 +1253,7 @@ requires-dist = [ { name = "numpy", specifier = ">=1.24" }, { name = "opencv-python" }, { name = "pre-commit", marker = "extra == 'dev'", specifier = ">=3.7" }, + { name = "pyarrow", specifier = ">=14.0" }, { name = "pytest", marker = "extra == 'dev'", specifier = ">=8.0" }, { name = "pytest-cov", marker = "extra == 'dev'", specifier = ">=5.0" }, { name = "pyyaml" },