Add data replay
This commit is contained in:
parent
d336b73aa3
commit
b5d46c4edd
2
.gitignore
vendored
2
.gitignore
vendored
@ -89,4 +89,4 @@ models/
|
||||
*.xvcd
|
||||
ufactory_usage/
|
||||
.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
|
||||
```
|
||||
|
||||
### 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 训练管道进行模仿学习训练。
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
|
||||
@ -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 = [
|
||||
"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"
|
||||
|
||||
@ -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)
|
||||
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:
|
||||
self.real_arm.set_mode(0)
|
||||
self.real_arm.set_state(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
|
||||
|
||||
|
||||
@ -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")
|
||||
|
||||
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.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" },
|
||||
|
||||
Loading…
Reference in New Issue
Block a user