Add data replay
This commit is contained in:
parent
d336b73aa3
commit
b5d46c4edd
2
.gitignore
vendored
2
.gitignore
vendored
@ -89,4 +89,4 @@ models/
|
|||||||
*.xvcd
|
*.xvcd
|
||||||
ufactory_usage/
|
ufactory_usage/
|
||||||
.history/
|
.history/
|
||||||
xarm7_manual_datas
|
datasets/
|
||||||
16
README_ZH.md
16
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
|
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训练
|
### 5. Lerobot训练
|
||||||
|
|
||||||
采集数据后,使用 LeRobot 训练管道进行模仿学习训练。
|
采集数据后,使用 LeRobot 训练管道进行模仿学习训练。
|
||||||
|
|||||||
@ -24,7 +24,7 @@ robot:
|
|||||||
fps: 30
|
fps: 30
|
||||||
|
|
||||||
dataset:
|
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"
|
repo_id: "ufactory/xarm7_manual_datas"
|
||||||
# Task description stored with each recorded frame.
|
# Task description stored with each recorded frame.
|
||||||
single_task: "Describe the task being demonstrated."
|
single_task: "Describe the task being demonstrated."
|
||||||
@ -37,3 +37,5 @@ dataset:
|
|||||||
# Store camera observations as videos.
|
# Store camera observations as videos.
|
||||||
video: true
|
video: true
|
||||||
push_to_hub: false
|
push_to_hub: false
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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
|
|
||||||
@ -25,6 +25,7 @@ classifiers = [
|
|||||||
]
|
]
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"numpy>=1.24",
|
"numpy>=1.24",
|
||||||
|
"pyarrow>=14.0",
|
||||||
"pyyaml",
|
"pyyaml",
|
||||||
"lerobot[intelrealsense]==0.4.3",
|
"lerobot[intelrealsense]==0.4.3",
|
||||||
"xarm-python-sdk",
|
"xarm-python-sdk",
|
||||||
@ -35,6 +36,7 @@ dependencies = [
|
|||||||
uf-robot-teleop = "lerobot_robot_ufactory.scripts.uf_robot_teleop:main"
|
uf-robot-teleop = "lerobot_robot_ufactory.scripts.uf_robot_teleop:main"
|
||||||
uf-lerobot-record = "lerobot_robot_ufactory.scripts.uf_lerobot_record: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-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-vive-calibrate = "lerobot_robot_ufactory.scripts.vive_calibrate:main"
|
||||||
uf-camera-view = "lerobot_robot_ufactory.scripts.uf_camera_view:main"
|
uf-camera-view = "lerobot_robot_ufactory.scripts.uf_camera_view:main"
|
||||||
uf-camera-test = "lerobot_robot_ufactory.scripts.uf_camera_test:main"
|
uf-camera-test = "lerobot_robot_ufactory.scripts.uf_camera_test:main"
|
||||||
|
|||||||
@ -223,8 +223,16 @@ class UFRobot(Robot, Thread):
|
|||||||
if self._initial_point is None:
|
if self._initial_point is None:
|
||||||
raise RuntimeError("xArm initial point has not been loaded")
|
raise RuntimeError("xArm initial point has not been loaded")
|
||||||
|
|
||||||
self.real_arm.set_mode(0)
|
# The controller requires motion to be enabled again after an
|
||||||
self.real_arm.set_state(0)
|
# 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(
|
code = self.real_arm.set_servo_angle(
|
||||||
angle=self._initial_point,
|
angle=self._initial_point,
|
||||||
speed=ROBOT_RESET_SPEED_DEG,
|
speed=ROBOT_RESET_SPEED_DEG,
|
||||||
@ -271,7 +279,9 @@ class UFRobot(Robot, Thread):
|
|||||||
return
|
return
|
||||||
|
|
||||||
if self._control_space == "joint":
|
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":
|
elif self._control_space == "cartesian":
|
||||||
self.real_arm.set_mode(7)
|
self.real_arm.set_mode(7)
|
||||||
else:
|
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]
|
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)
|
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:
|
def send_action(self, action: dict) -> np.ndarray:
|
||||||
if not self._is_connected:
|
if not self._is_connected:
|
||||||
raise ConnectionError()
|
raise ConnectionError()
|
||||||
@ -458,17 +482,43 @@ class UFRobot(Robot, Thread):
|
|||||||
for i in range(self._dof):
|
for i in range(self._dof):
|
||||||
cmd_list[i] = action[f"{self.prefix}J{i+1}.pos"]
|
cmd_list[i] = action[f"{self.prefix}J{i+1}.pos"]
|
||||||
|
|
||||||
# TODO: make mode 6 compatible with wait=True
|
if self.config.joint_command_mode == 1:
|
||||||
if wait_== False and self.real_arm.mode != 6:
|
# set_servo_angle_j is an absolute target command. It is the
|
||||||
self.real_arm.set_mode(6)
|
# SDK's high-frequency interface and executes only the latest
|
||||||
self.real_arm.set_state(0)
|
# 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)
|
time.sleep(0.1)
|
||||||
elif wait_ and self.real_arm.mode != 0:
|
elif wait_ and self.real_arm.mode != 0:
|
||||||
self.real_arm.set_mode(0)
|
code = self.real_arm.set_mode(0)
|
||||||
self.real_arm.set_state(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)
|
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?
|
elif self._control_space == "cartesian": # unit: mm?
|
||||||
lin_spd = self._max_linear_velocity
|
lin_spd = self._max_linear_velocity
|
||||||
|
|
||||||
|
|||||||
@ -20,6 +20,7 @@ class UFRobotConfig(RobotConfig):
|
|||||||
manual_mode: bool = False # xArm joint teaching mode; records state and optional gripper actions
|
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
|
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
|
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.
|
# start_joints and start_tcp_pose are intentionally disabled.
|
||||||
# Reset uses the xArm SDK initial_point instead of configuration poses.
|
# Reset uses the xArm SDK initial_point instead of configuration poses.
|
||||||
max_joint_velocity: int = 90 # °/s, only effective in joint control mode
|
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")
|
raise ValueError("teach_sensitivity must be between 1 and 5")
|
||||||
if self.manual_gripper_speed < 0:
|
if self.manual_gripper_speed < 0:
|
||||||
raise ValueError("manual_gripper_speed must be non-negative")
|
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")
|
||||||
|
|||||||
272
src/lerobot_robot_ufactory/scripts/uf_lerobot_replay.py
Normal file
272
src/lerobot_robot_ufactory/scripts/uf_lerobot_replay.py
Normal file
@ -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())
|
||||||
273
tests/test_lerobot_replay.py
Normal file
273
tests/test_lerobot_replay.py
Normal file
@ -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)
|
||||||
2
uv.lock
generated
2
uv.lock
generated
@ -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.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 = "numpy", version = "2.5.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" },
|
||||||
{ name = "opencv-python" },
|
{ name = "opencv-python" },
|
||||||
|
{ name = "pyarrow" },
|
||||||
{ name = "pyyaml" },
|
{ name = "pyyaml" },
|
||||||
{ name = "xarm-python-sdk" },
|
{ name = "xarm-python-sdk" },
|
||||||
]
|
]
|
||||||
@ -1252,6 +1253,7 @@ requires-dist = [
|
|||||||
{ name = "numpy", specifier = ">=1.24" },
|
{ name = "numpy", specifier = ">=1.24" },
|
||||||
{ name = "opencv-python" },
|
{ name = "opencv-python" },
|
||||||
{ name = "pre-commit", marker = "extra == 'dev'", specifier = ">=3.7" },
|
{ name = "pre-commit", marker = "extra == 'dev'", specifier = ">=3.7" },
|
||||||
|
{ name = "pyarrow", specifier = ">=14.0" },
|
||||||
{ name = "pytest", marker = "extra == 'dev'", specifier = ">=8.0" },
|
{ name = "pytest", marker = "extra == 'dev'", specifier = ">=8.0" },
|
||||||
{ name = "pytest-cov", marker = "extra == 'dev'", specifier = ">=5.0" },
|
{ name = "pytest-cov", marker = "extra == 'dev'", specifier = ">=5.0" },
|
||||||
{ name = "pyyaml" },
|
{ name = "pyyaml" },
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user