mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
Merge remote-tracking branch 'origin/main' into codex/episode-video-streaming-byte-cache
This commit is contained in:
@@ -478,18 +478,19 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
"""Return the number of frames in the selected episodes."""
|
"""Return the number of frames in the selected episodes."""
|
||||||
return self.num_frames
|
return self.num_frames
|
||||||
|
|
||||||
def __getitem__(self, idx) -> dict:
|
def __getitem__(self, idx: int | slice) -> dict | list[dict]:
|
||||||
"""Return a single frame by index, with all transforms applied.
|
"""Return one frame or a slice of frames, with all transforms applied.
|
||||||
|
|
||||||
Loads the frame from the underlying HF dataset, expands delta-timestamp
|
Loads the frame from the underlying HF dataset, expands delta-timestamp
|
||||||
windows, decodes video frames, and applies image transforms. Delegates
|
windows, decodes video frames, and applies image transforms. Delegates
|
||||||
the core logic to :meth:`DatasetReader.get_item`.
|
the core logic to :class:`DatasetReader`.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
idx: Index into the (possibly episode-filtered) dataset.
|
idx: Integer index or slice into the possibly episode-filtered dataset.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dict mapping feature names to their tensor values for this frame.
|
A frame dictionary for an integer index, or a list of frame
|
||||||
|
dictionaries for a slice.
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
RuntimeError: If the dataset is currently being recorded and
|
RuntimeError: If the dataset is currently being recorded and
|
||||||
@@ -499,6 +500,9 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"Cannot read from a dataset that is being recorded. Call finalize() first, then access items."
|
"Cannot read from a dataset that is being recorded. Call finalize() first, then access items."
|
||||||
)
|
)
|
||||||
|
if isinstance(idx, slice):
|
||||||
|
return [self[item_idx] for item_idx in range(*idx.indices(len(self)))]
|
||||||
|
|
||||||
reader = self._ensure_reader()
|
reader = self._ensure_reader()
|
||||||
if reader.hf_dataset is None:
|
if reader.hf_dataset is None:
|
||||||
# One-shot load after finalize()
|
# One-shot load after finalize()
|
||||||
|
|||||||
@@ -322,7 +322,7 @@ class HILSerlRobotEnvConfig(EnvConfig):
|
|||||||
class LiberoEnv(EnvConfig):
|
class LiberoEnv(EnvConfig):
|
||||||
task: str = "libero_10" # can also choose libero_spatial, libero_object, etc.
|
task: str = "libero_10" # can also choose libero_spatial, libero_object, etc.
|
||||||
task_ids: list[int] | None = None
|
task_ids: list[int] | None = None
|
||||||
fps: int = 30
|
fps: int = 20 # Must match robosuite's default control_freq (20 Hz)
|
||||||
episode_length: int | None = None
|
episode_length: int | None = None
|
||||||
obs_type: str = "pixels_agent_pos"
|
obs_type: str = "pixels_agent_pos"
|
||||||
render_mode: str = "rgb_array"
|
render_mode: str = "rgb_array"
|
||||||
@@ -354,6 +354,9 @@ class LiberoEnv(EnvConfig):
|
|||||||
control_mode: str = "relative" # or "absolute"
|
control_mode: str = "relative" # or "absolute"
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
|
if self.fps <= 0:
|
||||||
|
raise ValueError(f"fps must be positive, got {self.fps}")
|
||||||
|
|
||||||
if self.obs_type == "pixels":
|
if self.obs_type == "pixels":
|
||||||
self.features[LIBERO_KEY_PIXELS_AGENTVIEW] = PolicyFeature(
|
self.features[LIBERO_KEY_PIXELS_AGENTVIEW] = PolicyFeature(
|
||||||
type=FeatureType.VISUAL, shape=(self.observation_height, self.observation_width, 3)
|
type=FeatureType.VISUAL, shape=(self.observation_height, self.observation_width, 3)
|
||||||
@@ -412,6 +415,7 @@ class LiberoEnv(EnvConfig):
|
|||||||
"render_mode": self.render_mode,
|
"render_mode": self.render_mode,
|
||||||
"observation_height": self.observation_height,
|
"observation_height": self.observation_height,
|
||||||
"observation_width": self.observation_width,
|
"observation_width": self.observation_width,
|
||||||
|
"control_freq": self.fps,
|
||||||
}
|
}
|
||||||
if self.task_ids is not None:
|
if self.task_ids is not None:
|
||||||
kwargs["task_ids"] = self.task_ids
|
kwargs["task_ids"] = self.task_ids
|
||||||
|
|||||||
@@ -125,10 +125,13 @@ class LiberoEnv(gym.Env):
|
|||||||
n_envs: int = 1,
|
n_envs: int = 1,
|
||||||
camera_name_mapping: dict[str, str] | None = None,
|
camera_name_mapping: dict[str, str] | None = None,
|
||||||
num_steps_wait: int = 10,
|
num_steps_wait: int = 10,
|
||||||
|
control_freq: int = 20,
|
||||||
control_mode: str = "relative",
|
control_mode: str = "relative",
|
||||||
is_libero_plus: bool = False,
|
is_libero_plus: bool = False,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
if control_freq <= 0:
|
||||||
|
raise ValueError(f"control_freq must be positive, got {control_freq}")
|
||||||
self.task_id = task_id
|
self.task_id = task_id
|
||||||
self.is_libero_plus = is_libero_plus
|
self.is_libero_plus = is_libero_plus
|
||||||
self.obs_type = obs_type
|
self.obs_type = obs_type
|
||||||
@@ -154,6 +157,7 @@ class LiberoEnv(gym.Env):
|
|||||||
}
|
}
|
||||||
self.camera_name_mapping = camera_name_mapping
|
self.camera_name_mapping = camera_name_mapping
|
||||||
self.num_steps_wait = num_steps_wait
|
self.num_steps_wait = num_steps_wait
|
||||||
|
self.control_freq = control_freq
|
||||||
self.episode_index = episode_index
|
self.episode_index = episode_index
|
||||||
self.episode_length = episode_length
|
self.episode_length = episode_length
|
||||||
# Load once and keep
|
# Load once and keep
|
||||||
@@ -260,6 +264,7 @@ class LiberoEnv(gym.Env):
|
|||||||
bddl_file_name=self._task_bddl_file,
|
bddl_file_name=self._task_bddl_file,
|
||||||
camera_heights=self.observation_height,
|
camera_heights=self.observation_height,
|
||||||
camera_widths=self.observation_width,
|
camera_widths=self.observation_width,
|
||||||
|
control_freq=self.control_freq,
|
||||||
)
|
)
|
||||||
env.reset()
|
env.reset()
|
||||||
self._env = env
|
self._env = env
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ import logging
|
|||||||
import time
|
import time
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from functools import cached_property
|
|
||||||
from typing import TYPE_CHECKING, Any, TypedDict
|
from typing import TYPE_CHECKING, Any, TypedDict
|
||||||
|
|
||||||
from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected
|
from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected
|
||||||
@@ -854,7 +853,7 @@ class DamiaoMotorsBus(MotorsBusBase):
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Motor {motor_obj} doesn't have a valid recv_id (None).")
|
raise ValueError(f"Motor {motor_obj} doesn't have a valid recv_id (None).")
|
||||||
|
|
||||||
@cached_property
|
@property
|
||||||
def is_calibrated(self) -> bool:
|
def is_calibrated(self) -> bool:
|
||||||
"""Check if motors are calibrated."""
|
"""Check if motors are calibrated."""
|
||||||
return bool(self.calibration)
|
return bool(self.calibration)
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import abc
|
import abc
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
from collections.abc import Sequence
|
from collections.abc import Sequence
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
@@ -818,13 +819,13 @@ class SerialMotorsBus(MotorsBusBase):
|
|||||||
"""
|
"""
|
||||||
motor_names = self._get_motors_list(motors)
|
motor_names = self._get_motors_list(motors)
|
||||||
|
|
||||||
start_positions = self.sync_read("Present_Position", motor_names, normalize=False)
|
start_positions = self.sync_read("Present_Position", motor_names, normalize=False, num_retry=5)
|
||||||
mins = start_positions.copy()
|
mins = start_positions.copy()
|
||||||
maxes = start_positions.copy()
|
maxes = start_positions.copy()
|
||||||
|
|
||||||
user_pressed_enter = False
|
user_pressed_enter = False
|
||||||
while not user_pressed_enter:
|
while not user_pressed_enter:
|
||||||
positions = self.sync_read("Present_Position", motor_names, normalize=False)
|
positions = self.sync_read("Present_Position", motor_names, normalize=False, num_retry=5)
|
||||||
mins = {motor: min(positions[motor], min_) for motor, min_ in mins.items()}
|
mins = {motor: min(positions[motor], min_) for motor, min_ in mins.items()}
|
||||||
maxes = {motor: max(positions[motor], max_) for motor, max_ in maxes.items()}
|
maxes = {motor: max(positions[motor], max_) for motor, max_ in maxes.items()}
|
||||||
|
|
||||||
@@ -837,9 +838,12 @@ class SerialMotorsBus(MotorsBusBase):
|
|||||||
if enter_pressed():
|
if enter_pressed():
|
||||||
user_pressed_enter = True
|
user_pressed_enter = True
|
||||||
|
|
||||||
if display_values and not user_pressed_enter:
|
if not user_pressed_enter:
|
||||||
# Move cursor up to overwrite the previous output
|
if display_values:
|
||||||
move_cursor_up(len(motor_names) + 3)
|
# Move cursor up to overwrite the previous output
|
||||||
|
move_cursor_up(len(motor_names) + 3)
|
||||||
|
# Throttle reads even when the live table is disabled.
|
||||||
|
time.sleep(0.02)
|
||||||
|
|
||||||
same_min_max = [motor for motor in motor_names if mins[motor] == maxes[motor]]
|
same_min_max = [motor for motor in motor_names if mins[motor] == maxes[motor]]
|
||||||
if same_min_max:
|
if same_min_max:
|
||||||
|
|||||||
@@ -79,6 +79,8 @@ class DiffusionConfig(PreTrainedConfig):
|
|||||||
use_film_scale_modulation: FiLM (https://huggingface.co/papers/1709.07871) is used for the Unet conditioning.
|
use_film_scale_modulation: FiLM (https://huggingface.co/papers/1709.07871) is used for the Unet conditioning.
|
||||||
Bias modulation is used be default, while this parameter indicates whether to also use scale
|
Bias modulation is used be default, while this parameter indicates whether to also use scale
|
||||||
modulation.
|
modulation.
|
||||||
|
gradient_checkpointing: Whether to checkpoint the Unet residual blocks during training. This reduces
|
||||||
|
activation memory at the cost of recomputing those blocks during the backward pass.
|
||||||
noise_scheduler_type: Name of the noise scheduler to use. Supported options: ["DDPM", "DDIM"].
|
noise_scheduler_type: Name of the noise scheduler to use. Supported options: ["DDPM", "DDIM"].
|
||||||
num_train_timesteps: Number of diffusion steps for the forward diffusion schedule.
|
num_train_timesteps: Number of diffusion steps for the forward diffusion schedule.
|
||||||
beta_schedule: Name of the diffusion beta schedule as per DDPMScheduler from Hugging Face diffusers.
|
beta_schedule: Name of the diffusion beta schedule as per DDPMScheduler from Hugging Face diffusers.
|
||||||
@@ -132,6 +134,7 @@ class DiffusionConfig(PreTrainedConfig):
|
|||||||
n_groups: int = 8
|
n_groups: int = 8
|
||||||
diffusion_step_embed_dim: int = 128
|
diffusion_step_embed_dim: int = 128
|
||||||
use_film_scale_modulation: bool = True
|
use_film_scale_modulation: bool = True
|
||||||
|
gradient_checkpointing: bool = False
|
||||||
# Noise scheduler.
|
# Noise scheduler.
|
||||||
noise_scheduler_type: str = "DDPM"
|
noise_scheduler_type: str = "DDPM"
|
||||||
num_train_timesteps: int = 100
|
num_train_timesteps: int = 100
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ import torch
|
|||||||
import torch.nn.functional as F # noqa: N812
|
import torch.nn.functional as F # noqa: N812
|
||||||
import torchvision
|
import torchvision
|
||||||
from torch import Tensor, nn
|
from torch import Tensor, nn
|
||||||
|
from torch.utils.checkpoint import checkpoint
|
||||||
|
|
||||||
from lerobot.utils.constants import ACTION, OBS_ENV_STATE, OBS_IMAGES, OBS_STATE
|
from lerobot.utils.constants import ACTION, OBS_ENV_STATE, OBS_IMAGES, OBS_STATE
|
||||||
from lerobot.utils.import_utils import _diffusers_available, require_package
|
from lerobot.utils.import_utils import _diffusers_available, require_package
|
||||||
@@ -727,22 +728,35 @@ class DiffusionConditionalUnet1d(nn.Module):
|
|||||||
else:
|
else:
|
||||||
global_feature = timesteps_embed
|
global_feature = timesteps_embed
|
||||||
|
|
||||||
|
use_gc = self.config.gradient_checkpointing and self.training
|
||||||
|
|
||||||
# Run encoder, keeping track of skip features to pass to the decoder.
|
# Run encoder, keeping track of skip features to pass to the decoder.
|
||||||
encoder_skip_features: list[Tensor] = []
|
encoder_skip_features: list[Tensor] = []
|
||||||
for resnet, resnet2, downsample in self.down_modules:
|
for resnet, resnet2, downsample in self.down_modules:
|
||||||
x = resnet(x, global_feature)
|
if use_gc:
|
||||||
x = resnet2(x, global_feature)
|
x = checkpoint(resnet, x, global_feature, use_reentrant=False)
|
||||||
|
x = checkpoint(resnet2, x, global_feature, use_reentrant=False)
|
||||||
|
else:
|
||||||
|
x = resnet(x, global_feature)
|
||||||
|
x = resnet2(x, global_feature)
|
||||||
encoder_skip_features.append(x)
|
encoder_skip_features.append(x)
|
||||||
x = downsample(x)
|
x = downsample(x)
|
||||||
|
|
||||||
for mid_module in self.mid_modules:
|
for mid_module in self.mid_modules:
|
||||||
x = mid_module(x, global_feature)
|
if use_gc:
|
||||||
|
x = checkpoint(mid_module, x, global_feature, use_reentrant=False)
|
||||||
|
else:
|
||||||
|
x = mid_module(x, global_feature)
|
||||||
|
|
||||||
# Run decoder, using the skip features from the encoder.
|
# Run decoder, using the skip features from the encoder.
|
||||||
for resnet, resnet2, upsample in self.up_modules:
|
for resnet, resnet2, upsample in self.up_modules:
|
||||||
x = torch.cat((x, encoder_skip_features.pop()), dim=1)
|
x = torch.cat((x, encoder_skip_features.pop()), dim=1)
|
||||||
x = resnet(x, global_feature)
|
if use_gc:
|
||||||
x = resnet2(x, global_feature)
|
x = checkpoint(resnet, x, global_feature, use_reentrant=False)
|
||||||
|
x = checkpoint(resnet2, x, global_feature, use_reentrant=False)
|
||||||
|
else:
|
||||||
|
x = resnet(x, global_feature)
|
||||||
|
x = resnet2(x, global_feature)
|
||||||
x = upsample(x)
|
x = upsample(x)
|
||||||
|
|
||||||
x = self.final_conv(x)
|
x = self.final_conv(x)
|
||||||
|
|||||||
@@ -150,9 +150,6 @@ class OpenArmFollower(Robot):
|
|||||||
|
|
||||||
self.configure()
|
self.configure()
|
||||||
|
|
||||||
if self.is_calibrated:
|
|
||||||
self.bus.set_zero_position()
|
|
||||||
|
|
||||||
self.bus.enable_torque()
|
self.bus.enable_torque()
|
||||||
|
|
||||||
logger.info(f"{self} connected.")
|
logger.info(f"{self} connected.")
|
||||||
|
|||||||
@@ -51,19 +51,7 @@ from lerobot.teleoperators import ( # noqa: F401
|
|||||||
rebot_102_leader,
|
rebot_102_leader,
|
||||||
so_leader,
|
so_leader,
|
||||||
)
|
)
|
||||||
|
from lerobot.utils.import_utils import register_third_party_plugins
|
||||||
COMPATIBLE_DEVICES = [
|
|
||||||
"koch_follower",
|
|
||||||
"koch_leader",
|
|
||||||
"omx_follower",
|
|
||||||
"omx_leader",
|
|
||||||
"openarm_mini",
|
|
||||||
"so100_follower",
|
|
||||||
"so100_leader",
|
|
||||||
"so101_follower",
|
|
||||||
"so101_leader",
|
|
||||||
"lekiwi",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -80,18 +68,19 @@ class SetupConfig:
|
|||||||
|
|
||||||
@draccus.wrap()
|
@draccus.wrap()
|
||||||
def setup_motors(cfg: SetupConfig):
|
def setup_motors(cfg: SetupConfig):
|
||||||
if cfg.device.type not in COMPATIBLE_DEVICES:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
if isinstance(cfg.device, RobotConfig):
|
if isinstance(cfg.device, RobotConfig):
|
||||||
device = make_robot_from_config(cfg.device)
|
device = make_robot_from_config(cfg.device)
|
||||||
else:
|
else:
|
||||||
device = make_teleoperator_from_config(cfg.device)
|
device = make_teleoperator_from_config(cfg.device)
|
||||||
|
|
||||||
device.setup_motors()
|
setup = getattr(device, "setup_motors", None)
|
||||||
|
if not callable(setup):
|
||||||
|
raise NotImplementedError(f"Device type '{cfg.device.type}' does not support motor setup.")
|
||||||
|
setup()
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
|
register_third_party_plugins()
|
||||||
setup_motors()
|
setup_motors()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -23,3 +23,5 @@ from ..config import TeleoperatorConfig
|
|||||||
@dataclass
|
@dataclass
|
||||||
class GamepadTeleopConfig(TeleoperatorConfig):
|
class GamepadTeleopConfig(TeleoperatorConfig):
|
||||||
use_gripper: bool = True
|
use_gripper: bool = True
|
||||||
|
# Use hidapi instead of pygame for controllers that pygame cannot detect reliably.
|
||||||
|
hidapi_fallback: bool = False
|
||||||
|
|||||||
@@ -14,6 +14,7 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
|
import logging
|
||||||
import sys
|
import sys
|
||||||
from enum import IntEnum
|
from enum import IntEnum
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -27,6 +28,8 @@ from ..teleoperator import Teleoperator
|
|||||||
from ..utils import TeleopEvents
|
from ..utils import TeleopEvents
|
||||||
from .configuration_gamepad import GamepadTeleopConfig
|
from .configuration_gamepad import GamepadTeleopConfig
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class GripperAction(IntEnum):
|
class GripperAction(IntEnum):
|
||||||
CLOSE = 0
|
CLOSE = 0
|
||||||
@@ -56,6 +59,13 @@ class GamepadTeleop(Teleoperator):
|
|||||||
|
|
||||||
self.gamepad = None
|
self.gamepad = None
|
||||||
|
|
||||||
|
self.hidapi_fallback = config.hidapi_fallback
|
||||||
|
if sys.platform == "darwin" and not self.hidapi_fallback:
|
||||||
|
logger.warning(
|
||||||
|
"On macOS, pygame may not reliably detect input from some controllers. "
|
||||||
|
"If you experience issues, set `hidapi_fallback=true`."
|
||||||
|
)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def action_features(self) -> dict:
|
def action_features(self) -> dict:
|
||||||
if self.config.use_gripper:
|
if self.config.use_gripper:
|
||||||
@@ -76,9 +86,7 @@ class GamepadTeleop(Teleoperator):
|
|||||||
return {}
|
return {}
|
||||||
|
|
||||||
def connect(self) -> None:
|
def connect(self) -> None:
|
||||||
# use HidApi for macos
|
if self.hidapi_fallback:
|
||||||
if sys.platform == "darwin":
|
|
||||||
# NOTE: On macOS, pygame doesn’t reliably detect input from some controllers so we fall back to hidapi
|
|
||||||
from .gamepad_utils import GamepadControllerHID as Gamepad
|
from .gamepad_utils import GamepadControllerHID as Gamepad
|
||||||
else:
|
else:
|
||||||
from .gamepad_utils import GamepadController as Gamepad
|
from .gamepad_utils import GamepadController as Gamepad
|
||||||
|
|||||||
@@ -114,6 +114,20 @@ def test_dataset_initialization(tmp_path, lerobot_dataset_factory):
|
|||||||
assert dataset.num_frames == len(dataset)
|
assert dataset.num_frames == len(dataset)
|
||||||
|
|
||||||
|
|
||||||
|
def test_dataset_slice(tmp_path, lerobot_dataset_factory):
|
||||||
|
dataset = lerobot_dataset_factory(
|
||||||
|
root=tmp_path / "test", total_episodes=3, total_frames=30, use_videos=False
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(dataset[:5]) == 5
|
||||||
|
assert len(dataset[::2]) == (len(dataset) + 1) // 2
|
||||||
|
assert [item["index"].item() for item in dataset[4::-1]] == [4, 3, 2, 1, 0]
|
||||||
|
assert [item["index"].item() for item in dataset[-3:]] == list(range(len(dataset) - 3, len(dataset)))
|
||||||
|
assert dataset[len(dataset) :] == []
|
||||||
|
assert isinstance(dataset[0], dict)
|
||||||
|
assert dataset[:1][0].keys() == dataset[0].keys()
|
||||||
|
|
||||||
|
|
||||||
# TODO(rcadene, aliberts): do not run LeRobotDataset.create, instead refactor LeRobotDatasetMetadata.create
|
# TODO(rcadene, aliberts): do not run LeRobotDataset.create, instead refactor LeRobotDatasetMetadata.create
|
||||||
# and test the small resulting function that validates the features
|
# and test the small resulting function that validates the features
|
||||||
def test_dataset_feature_with_forward_slash_raises_error():
|
def test_dataset_feature_with_forward_slash_raises_error():
|
||||||
@@ -1741,6 +1755,38 @@ def test_delta_timestamps_query_returns_correct_values(tmp_path, empty_lerobot_d
|
|||||||
assert is_pad == [True, False], f"Expected [True, False], got {is_pad}"
|
assert is_pad == [True, False], f"Expected [True, False], got {is_pad}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_dataset_slice_with_delta_timestamps(tmp_path, empty_lerobot_dataset_factory):
|
||||||
|
features = {
|
||||||
|
"observation.state": {"dtype": "float32", "shape": (1,), "names": ["x"]},
|
||||||
|
}
|
||||||
|
dataset = empty_lerobot_dataset_factory(
|
||||||
|
root=tmp_path / "test_slice_delta", features=features, use_videos=False, fps=10
|
||||||
|
)
|
||||||
|
|
||||||
|
for frame_idx in range(5):
|
||||||
|
dataset.add_frame(
|
||||||
|
{
|
||||||
|
"observation.state": torch.tensor([frame_idx], dtype=torch.float32),
|
||||||
|
"task": "task_0",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
dataset.save_episode()
|
||||||
|
dataset.finalize()
|
||||||
|
|
||||||
|
sliced_dataset = LeRobotDataset(
|
||||||
|
dataset.repo_id,
|
||||||
|
root=dataset.root,
|
||||||
|
delta_timestamps={"observation.state": [-0.1, 0.0]},
|
||||||
|
tolerance_s=0.04,
|
||||||
|
)
|
||||||
|
|
||||||
|
items = sliced_dataset[:2]
|
||||||
|
|
||||||
|
assert items[0]["observation.state"].tolist() == [0.0, 0.0]
|
||||||
|
assert items[0]["observation.state_is_pad"].tolist() == [True, False]
|
||||||
|
assert items[1]["observation.state"].tolist() == [0.0, 1.0]
|
||||||
|
|
||||||
|
|
||||||
def test_episode_filter_filters_dataset(tmp_path, lerobot_dataset_factory):
|
def test_episode_filter_filters_dataset(tmp_path, lerobot_dataset_factory):
|
||||||
"""episode_filter on LeRobotDataset narrows the loaded dataset to matching episodes."""
|
"""episode_filter on LeRobotDataset narrows the loaded dataset to matching episodes."""
|
||||||
dataset = lerobot_dataset_factory(root=tmp_path / "test", total_episodes=8, total_frames=200)
|
dataset = lerobot_dataset_factory(root=tmp_path / "test", total_episodes=8, total_frames=200)
|
||||||
|
|||||||
@@ -35,6 +35,17 @@ def test_unknown_type():
|
|||||||
make_env_config("nonexistent")
|
make_env_config("nonexistent")
|
||||||
|
|
||||||
|
|
||||||
|
def test_libero_fps_controls_simulator_frequency():
|
||||||
|
cfg = LiberoEnv(fps=17)
|
||||||
|
|
||||||
|
assert cfg.gym_kwargs["control_freq"] == 17
|
||||||
|
|
||||||
|
|
||||||
|
def test_libero_rejects_nonpositive_fps():
|
||||||
|
with pytest.raises(ValueError, match="fps must be positive"):
|
||||||
|
LiberoEnv(fps=0)
|
||||||
|
|
||||||
|
|
||||||
def test_identity_processors():
|
def test_identity_processors():
|
||||||
"""Base class get_env_processors() returns identity pipelines."""
|
"""Base class get_env_processors() returns identity pipelines."""
|
||||||
cfg = make_env_config("aloha")
|
cfg = make_env_config("aloha")
|
||||||
|
|||||||
@@ -405,12 +405,18 @@ def test_record_ranges_of_motion(mock_motors, dummy_motors):
|
|||||||
read_pos_stub = mock_motors.build_sequential_sync_read_stub(
|
read_pos_stub = mock_motors.build_sequential_sync_read_stub(
|
||||||
*X_SERIES_CONTROL_TABLE["Present_Position"], positions
|
*X_SERIES_CONTROL_TABLE["Present_Position"], positions
|
||||||
)
|
)
|
||||||
with patch("lerobot.motors.motors_bus.enter_pressed", side_effect=[False, True]):
|
bus = DynamixelMotorsBus(port=mock_motors.port, motors=dummy_motors)
|
||||||
bus = DynamixelMotorsBus(port=mock_motors.port, motors=dummy_motors)
|
bus.connect(handshake=False)
|
||||||
bus.connect(handshake=False)
|
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("lerobot.motors.motors_bus.enter_pressed", side_effect=[False, True]),
|
||||||
|
patch("lerobot.motors.motors_bus.time.sleep") as mock_sleep,
|
||||||
|
patch.object(bus, "sync_read", wraps=bus.sync_read) as mock_sync_read,
|
||||||
|
):
|
||||||
mins, maxes = bus.record_ranges_of_motion(display_values=False)
|
mins, maxes = bus.record_ranges_of_motion(display_values=False)
|
||||||
|
|
||||||
assert mock_motors.stubs[read_pos_stub].calls == 3
|
assert mock_motors.stubs[read_pos_stub].calls == 3
|
||||||
|
assert all(call.kwargs["num_retry"] == 5 for call in mock_sync_read.call_args_list)
|
||||||
|
mock_sleep.assert_called_once_with(0.02)
|
||||||
assert mins == expected_mins
|
assert mins == expected_mins
|
||||||
assert maxes == expected_maxes
|
assert maxes == expected_maxes
|
||||||
|
|||||||
@@ -509,12 +509,18 @@ def test_record_ranges_of_motion(mock_motors, dummy_motors):
|
|||||||
stub = mock_motors.build_sequential_sync_read_stub(
|
stub = mock_motors.build_sequential_sync_read_stub(
|
||||||
*STS_SMS_SERIES_CONTROL_TABLE["Present_Position"], positions
|
*STS_SMS_SERIES_CONTROL_TABLE["Present_Position"], positions
|
||||||
)
|
)
|
||||||
with patch("lerobot.motors.motors_bus.enter_pressed", side_effect=[False, True]):
|
bus = FeetechMotorsBus(port=mock_motors.port, motors=dummy_motors)
|
||||||
bus = FeetechMotorsBus(port=mock_motors.port, motors=dummy_motors)
|
bus.connect(handshake=False)
|
||||||
bus.connect(handshake=False)
|
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("lerobot.motors.motors_bus.enter_pressed", side_effect=[False, True]),
|
||||||
|
patch("lerobot.motors.motors_bus.time.sleep") as mock_sleep,
|
||||||
|
patch.object(bus, "sync_read", wraps=bus.sync_read) as mock_sync_read,
|
||||||
|
):
|
||||||
mins, maxes = bus.record_ranges_of_motion(display_values=False)
|
mins, maxes = bus.record_ranges_of_motion(display_values=False)
|
||||||
|
|
||||||
assert mock_motors.stubs[stub].calls == 3
|
assert mock_motors.stubs[stub].calls == 3
|
||||||
|
assert all(call.kwargs["num_retry"] == 5 for call in mock_sync_read.call_args_list)
|
||||||
|
mock_sleep.assert_called_once_with(0.02)
|
||||||
assert mins == expected_mins
|
assert mins == expected_mins
|
||||||
assert maxes == expected_maxes
|
assert maxes == expected_maxes
|
||||||
|
|||||||
@@ -0,0 +1,49 @@
|
|||||||
|
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
import lerobot.scripts.lerobot_setup_motors as motors_module
|
||||||
|
|
||||||
|
|
||||||
|
def test_main_registers_plugins_before_parsing(monkeypatch):
|
||||||
|
calls = []
|
||||||
|
monkeypatch.setattr(motors_module, "register_third_party_plugins", lambda: calls.append("register"))
|
||||||
|
monkeypatch.setattr(motors_module, "setup_motors", lambda: calls.append("setup"))
|
||||||
|
|
||||||
|
motors_module.main()
|
||||||
|
|
||||||
|
assert calls == ["register", "setup"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_setup_motors_accepts_third_party_device(monkeypatch):
|
||||||
|
device = MagicMock()
|
||||||
|
monkeypatch.setattr(motors_module, "make_teleoperator_from_config", lambda _: device)
|
||||||
|
cfg = SimpleNamespace(device=SimpleNamespace(type="third_party"))
|
||||||
|
|
||||||
|
motors_module.setup_motors.__wrapped__(cfg)
|
||||||
|
|
||||||
|
device.setup_motors.assert_called_once_with()
|
||||||
|
|
||||||
|
|
||||||
|
def test_setup_motors_reports_unsupported_device(monkeypatch):
|
||||||
|
device = object()
|
||||||
|
monkeypatch.setattr(motors_module, "make_teleoperator_from_config", lambda _: device)
|
||||||
|
cfg = SimpleNamespace(device=SimpleNamespace(type="third_party"))
|
||||||
|
|
||||||
|
with pytest.raises(NotImplementedError, match="third_party"):
|
||||||
|
motors_module.setup_motors.__wrapped__(cfg)
|
||||||
Reference in New Issue
Block a user