diff --git a/pyproject.toml b/pyproject.toml index 0e9fba586..4b50a77f8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -108,9 +108,9 @@ training = [ "wandb>=0.24.0,<0.25.0", ] hardware = [ - "pynput>=1.7.8,<1.9.0", - "pyserial>=3.5,<4.0", - "deepdiff>=7.0.1,<9.0.0", + "lerobot[pynput-dep]", + "lerobot[pyserial-dep]", + "lerobot[deepdiff-dep]", ] viz = [ "rerun-sdk>=0.24.0,<0.27.0", @@ -136,10 +136,14 @@ scipy-dep = ["scipy>=1.14.0,<2.0.0"] diffusers-dep = ["diffusers>=0.27.2,<0.36.0"] qwen-vl-utils-dep = ["qwen-vl-utils>=0.0.11,<0.1.0"] matplotlib-dep = ["matplotlib>=3.10.3,<4.0.0", "contourpy>=1.3.0,<2.0.0"] # NOTE: Explicitly listing contourpy helps the resolver converge faster. +pyserial-dep = ["pyserial>=3.5,<4.0"] +deepdiff-dep = ["deepdiff>=7.0.1,<9.0.0"] +pynput-dep = ["pynput>=1.7.8,<1.9.0"] +pyzmq-dep = ["pyzmq>=26.2.1,<28.0.0"] # Motors -feetech = ["feetech-servo-sdk>=1.0.0,<2.0.0"] -dynamixel = ["dynamixel-sdk>=3.7.31,<3.9.0"] +feetech = ["feetech-servo-sdk>=1.0.0,<2.0.0", "lerobot[pyserial-dep]", "lerobot[deepdiff-dep]"] +dynamixel = ["dynamixel-sdk>=3.7.31,<3.9.0", "lerobot[pyserial-dep]", "lerobot[deepdiff-dep]"] damiao = ["lerobot[can-dep]"] robstride = ["lerobot[can-dep]"] @@ -147,10 +151,11 @@ robstride = ["lerobot[can-dep]"] openarms = ["lerobot[damiao]"] gamepad = ["lerobot[pygame-dep]", "hidapi>=0.14.0,<0.15.0"] hopejr = ["lerobot[feetech]", "lerobot[pygame-dep]"] -lekiwi = ["lerobot[feetech]", "pyzmq>=26.2.1,<28.0.0"] +lekiwi = ["lerobot[feetech]", "lerobot[pyzmq-dep]"] unitree_g1 = [ # "unitree-sdk2==1.0.1", - "pyzmq>=26.2.1,<28.0.0", + "lerobot[pyzmq-dep]", + "lerobot[pyserial-dep]", "onnxruntime>=1.16.0,<2.0.0", "onnx>=1.16.0,<2.0.0", "meshcat>=0.3.0,<0.4.0", diff --git a/src/lerobot/cameras/reachy2_camera/reachy2_camera.py b/src/lerobot/cameras/reachy2_camera/reachy2_camera.py index 9bef957bc..9b7e7a2e0 100644 --- a/src/lerobot/cameras/reachy2_camera/reachy2_camera.py +++ b/src/lerobot/cameras/reachy2_camera/reachy2_camera.py @@ -33,7 +33,7 @@ import cv2 # type: ignore # TODO: add type stubs for OpenCV import numpy as np # type: ignore # TODO: add type stubs for numpy from lerobot.utils.decorators import check_if_not_connected -from lerobot.utils.import_utils import _reachy2_sdk_available +from lerobot.utils.import_utils import _reachy2_sdk_available, require_package if TYPE_CHECKING or _reachy2_sdk_available: from reachy2_sdk.media.camera import CameraView @@ -76,6 +76,7 @@ class Reachy2Camera(Camera): Args: config: The configuration settings for the camera. """ + require_package("reachy2_sdk", extra="reachy2") super().__init__(config) self.config = config diff --git a/src/lerobot/cameras/realsense/camera_realsense.py b/src/lerobot/cameras/realsense/camera_realsense.py index d80ec8093..6363cc0bc 100644 --- a/src/lerobot/cameras/realsense/camera_realsense.py +++ b/src/lerobot/cameras/realsense/camera_realsense.py @@ -19,16 +19,18 @@ Provides the RealSenseCamera class for capturing frames from Intel RealSense cam import logging import time from threading import Event, Lock, Thread -from typing import Any +from typing import TYPE_CHECKING, Any import cv2 # type: ignore # TODO: add type stubs for OpenCV import numpy as np # type: ignore # TODO: add type stubs for numpy from numpy.typing import NDArray # type: ignore # TODO: add type stubs for numpy.typing -try: - import pyrealsense2 as rs # type: ignore # TODO: add type stubs for pyrealsense2 -except Exception as e: - logging.info(f"Could not import realsense: {e}") +from lerobot.utils.import_utils import _pyrealsense2_available, require_package + +if TYPE_CHECKING or _pyrealsense2_available: + import pyrealsense2 as rs +else: + rs = None from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected from lerobot.utils.errors import DeviceNotConnectedError @@ -112,7 +114,7 @@ class RealSenseCamera(Camera): Args: config: The configuration settings for the camera. """ - + require_package("pyrealsense2", extra="intelrealsense") super().__init__(config) self.config = config diff --git a/src/lerobot/cameras/zmq/camera_zmq.py b/src/lerobot/cameras/zmq/camera_zmq.py index 2fbe50d8b..1b0be5de6 100644 --- a/src/lerobot/cameras/zmq/camera_zmq.py +++ b/src/lerobot/cameras/zmq/camera_zmq.py @@ -28,12 +28,19 @@ import json import logging import time from threading import Event, Lock, Thread -from typing import Any +from typing import TYPE_CHECKING, Any import cv2 import numpy as np from numpy.typing import NDArray +from lerobot.utils.import_utils import _zmq_available, require_package + +if TYPE_CHECKING or _zmq_available: + import zmq +else: + zmq = None + from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected from lerobot.utils.errors import DeviceNotConnectedError @@ -74,8 +81,8 @@ class ZMQCamera(Camera): """ def __init__(self, config: ZMQCameraConfig): + require_package("pyzmq", extra="pyzmq-dep", import_name="zmq") super().__init__(config) - import zmq self.config = config self.server_address = config.server_address @@ -117,8 +124,6 @@ class ZMQCamera(Camera): logger.info(f"Connecting to {self}...") try: - import zmq - self.context = zmq.Context() self.socket = self.context.socket(zmq.SUB) self.socket.setsockopt_string(zmq.SUBSCRIBE, "") @@ -180,11 +185,8 @@ class ZMQCamera(Camera): try: message = self.socket.recv_string() - except Exception as e: - # zmq is lazy-imported in connect(), so check by name to avoid a top-level import - if type(e).__name__ == "Again": - raise TimeoutError(f"{self} timeout after {self.timeout_ms}ms") from e - raise + except zmq.Again as e: + raise TimeoutError(f"{self} timeout after {self.timeout_ms}ms") from e # Decode JSON message data = json.loads(message) diff --git a/src/lerobot/common/control_utils.py b/src/lerobot/common/control_utils.py index 530955078..efbcc082f 100644 --- a/src/lerobot/common/control_utils.py +++ b/src/lerobot/common/control_utils.py @@ -28,6 +28,12 @@ import numpy as np import torch from lerobot.policies import PreTrainedPolicy, prepare_observation_for_inference +from lerobot.utils.import_utils import _deepdiff_available, require_package + +if TYPE_CHECKING or _deepdiff_available: + from deepdiff import DeepDiff +else: + DeepDiff = None if TYPE_CHECKING: from lerobot.datasets import LeRobotDataset @@ -217,10 +223,7 @@ def sanity_check_dataset_robot_compatibility( Raises: ValueError: If any of the checked metadata fields do not match. """ - from lerobot.utils.import_utils import require_package - - require_package("deepdiff", extra="hardware") - from deepdiff import DeepDiff + require_package("deepdiff", extra="deepdiff-dep") from lerobot.utils.constants import DEFAULT_FEATURES diff --git a/src/lerobot/datasets/image_writer.py b/src/lerobot/datasets/image_writer.py index 603067757..8fb5804a5 100644 --- a/src/lerobot/datasets/image_writer.py +++ b/src/lerobot/datasets/image_writer.py @@ -30,13 +30,13 @@ def safe_stop_image_writer(func): def wrapper(*args, **kwargs): try: return func(*args, **kwargs) - except Exception as e: + except BaseException: dataset = kwargs.get("dataset") writer = getattr(dataset, "writer", None) if dataset else None if writer is not None and writer.image_writer is not None: logger.warning("Waiting for image writer to terminate...") writer.image_writer.stop() - raise e + raise return wrapper diff --git a/src/lerobot/model/kinematics.py b/src/lerobot/model/kinematics.py index 95d3b235c..01705ded5 100644 --- a/src/lerobot/model/kinematics.py +++ b/src/lerobot/model/kinematics.py @@ -12,8 +12,19 @@ # See the License for the specific language governing permissions and # limitations under the License. +from __future__ import annotations + +from typing import TYPE_CHECKING + import numpy as np +from lerobot.utils.import_utils import _placo_available, require_package + +if TYPE_CHECKING or _placo_available: + import placo # type: ignore[import-not-found] +else: + placo = None + class RobotKinematics: """Robot kinematics using placo library for forward and inverse kinematics.""" @@ -32,13 +43,7 @@ class RobotKinematics: target_frame_name (str): Name of the end-effector frame in the URDF joint_names (list[str] | None): List of joint names to use for the kinematics solver """ - try: - import placo # type: ignore[import-not-found] # C++ library with Python bindings, no type stubs available. TODO: Create stub file or request upstream typing support. - except ImportError as e: - raise ImportError( - "placo is required for RobotKinematics. " - "Please install the optional dependencies of `kinematics` in the package." - ) from e + require_package("placo", extra="placo-dep") self.robot = placo.RobotWrapper(urdf_path) self.solver = placo.KinematicsSolver(self.robot) diff --git a/src/lerobot/motors/damiao/damiao.py b/src/lerobot/motors/damiao/damiao.py index ae619f159..572741cb4 100644 --- a/src/lerobot/motors/damiao/damiao.py +++ b/src/lerobot/motors/damiao/damiao.py @@ -24,7 +24,7 @@ from functools import cached_property from typing import TYPE_CHECKING, Any, TypedDict from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected -from lerobot.utils.import_utils import _can_available +from lerobot.utils.import_utils import _can_available, require_package if TYPE_CHECKING or _can_available: import can @@ -111,6 +111,7 @@ class DamiaoMotorsBus(MotorsBusBase): bitrate: Nominal bitrate in bps (default: 1000000 = 1 Mbps) data_bitrate: Data bitrate for CAN FD in bps (default: 5000000 = 5 Mbps), ignored if use_can_fd is False """ + require_package("python-can", extra="damiao", import_name="can") super().__init__(port, motors, calibration) self.port = port self.can_interface = can_interface diff --git a/src/lerobot/motors/motors_bus.py b/src/lerobot/motors/motors_bus.py index 209489bb9..4688eaa7f 100644 --- a/src/lerobot/motors/motors_bus.py +++ b/src/lerobot/motors/motors_bus.py @@ -356,8 +356,8 @@ class SerialMotorsBus(MotorsBusBase): motors: dict[str, Motor], calibration: dict[str, MotorCalibration] | None = None, ): - require_package("pyserial", extra="hardware", import_name="serial") - require_package("deepdiff", extra="hardware") + require_package("pyserial", extra="pyserial-dep", import_name="serial") + require_package("deepdiff", extra="deepdiff-dep") super().__init__(port, motors, calibration) self.port_handler: PortHandler diff --git a/src/lerobot/motors/robstride/robstride.py b/src/lerobot/motors/robstride/robstride.py index f47e41509..ecde01e9a 100644 --- a/src/lerobot/motors/robstride/robstride.py +++ b/src/lerobot/motors/robstride/robstride.py @@ -23,12 +23,12 @@ from types import SimpleNamespace from typing import TYPE_CHECKING, Any, TypedDict from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected -from lerobot.utils.import_utils import _can_available +from lerobot.utils.import_utils import _can_available, require_package if TYPE_CHECKING or _can_available: import can else: - can = SimpleNamespace(Message=object, interface=None) + can = SimpleNamespace(Message=object, interface=None, BusABC=object) import numpy as np from lerobot.utils.errors import DeviceNotConnectedError @@ -106,6 +106,7 @@ class RobstrideMotorsBus(MotorsBusBase): bitrate: Nominal bitrate in bps (default: 1000000 = 1 Mbps) data_bitrate: Data bitrate for CAN FD in bps (default: 5000000 = 5 Mbps), ignored if use_can_fd is False """ + require_package("python-can", extra="robstride", import_name="can") super().__init__(port, motors, calibration) self.port = port self.can_interface = can_interface diff --git a/src/lerobot/optim/schedulers.py b/src/lerobot/optim/schedulers.py index 914edd2db..250650089 100644 --- a/src/lerobot/optim/schedulers.py +++ b/src/lerobot/optim/schedulers.py @@ -18,14 +18,21 @@ import logging import math from dataclasses import asdict, dataclass from pathlib import Path +from typing import TYPE_CHECKING import draccus from torch.optim import Optimizer from torch.optim.lr_scheduler import LambdaLR, LRScheduler from lerobot.utils.constants import SCHEDULER_STATE +from lerobot.utils.import_utils import _diffusers_available, require_package from lerobot.utils.io_utils import deserialize_json_into_object, write_json +if TYPE_CHECKING or _diffusers_available: + from diffusers.optimization import get_scheduler +else: + get_scheduler = None + @dataclass class LRSchedulerConfig(draccus.ChoiceRegistry, abc.ABC): @@ -47,10 +54,7 @@ class DiffuserSchedulerConfig(LRSchedulerConfig): num_warmup_steps: int | None = None def build(self, optimizer: Optimizer, num_training_steps: int) -> LambdaLR: - from lerobot.utils.import_utils import require_package - require_package("diffusers", extra="diffusion") - from diffusers.optimization import get_scheduler kwargs = {**asdict(self), "num_training_steps": num_training_steps, "optimizer": optimizer} return get_scheduler(**kwargs) diff --git a/src/lerobot/policies/diffusion/modeling_diffusion.py b/src/lerobot/policies/diffusion/modeling_diffusion.py index 5b3b97571..03203ffc8 100644 --- a/src/lerobot/policies/diffusion/modeling_diffusion.py +++ b/src/lerobot/policies/diffusion/modeling_diffusion.py @@ -23,6 +23,7 @@ TODO(alexander-soare): import math from collections import deque from collections.abc import Callable +from typing import TYPE_CHECKING import einops import numpy as np @@ -32,6 +33,14 @@ import torchvision from torch import Tensor, nn from lerobot.utils.constants import ACTION, OBS_ENV_STATE, OBS_IMAGES, OBS_STATE +from lerobot.utils.import_utils import _diffusers_available, require_package + +if TYPE_CHECKING or _diffusers_available: + from diffusers.schedulers.scheduling_ddim import DDIMScheduler + from diffusers.schedulers.scheduling_ddpm import DDPMScheduler +else: + DDIMScheduler = None + DDPMScheduler = None from ..pretrained import PreTrainedPolicy from ..utils import ( @@ -64,6 +73,7 @@ class DiffusionPolicy(PreTrainedPolicy): dataset_stats: Dataset statistics to be used for normalization. If not passed here, it is expected that they will be passed with a call to `load_state_dict` before the policy is used. """ + require_package("diffusers", extra="diffusion") super().__init__(config) config.validate_features() self.config = config @@ -155,11 +165,7 @@ def _make_noise_scheduler(name: str, **kwargs: dict): Factory for noise scheduler instances of the requested type. All kwargs are passed to the scheduler. """ - from lerobot.utils.import_utils import require_package - require_package("diffusers", extra="diffusion") - from diffusers.schedulers.scheduling_ddim import DDIMScheduler - from diffusers.schedulers.scheduling_ddpm import DDPMScheduler if name == "DDPM": return DDPMScheduler(**kwargs) diff --git a/src/lerobot/policies/groot/action_head/flow_matching_action_head.py b/src/lerobot/policies/groot/action_head/flow_matching_action_head.py index 4fda21ca5..2c1ca6014 100644 --- a/src/lerobot/policies/groot/action_head/flow_matching_action_head.py +++ b/src/lerobot/policies/groot/action_head/flow_matching_action_head.py @@ -204,7 +204,9 @@ class FlowmatchingActionHead(nn.Module): self.position_embedding = nn.Embedding(config.max_seq_len, self.input_embedding_dim) nn.init.normal_(self.position_embedding.weight, mean=0.0, std=0.02) - self.beta_dist = Beta(config.noise_beta_alpha, config.noise_beta_beta) + self._noise_beta_alpha = config.noise_beta_alpha + self._noise_beta_beta = config.noise_beta_beta + self._beta_dist = None self.num_timestep_buckets = config.num_timestep_buckets self.config = config self.set_trainable_parameters(config.tune_projector, config.tune_diffusion_model) @@ -249,7 +251,9 @@ class FlowmatchingActionHead(nn.Module): self.model.eval() def sample_time(self, batch_size, device, dtype): - sample = self.beta_dist.sample([batch_size]).to(device, dtype=dtype) + if self._beta_dist is None: + self._beta_dist = Beta(self._noise_beta_alpha, self._noise_beta_beta, validate_args=False) + sample = self._beta_dist.sample([batch_size]).to(device, dtype=dtype) return (self.config.noise_s - sample) / self.config.noise_s def prepare_input(self, batch: dict) -> BatchFeature: diff --git a/src/lerobot/policies/groot/eagle2_hg_model/processing_eagle2_5_vl.py b/src/lerobot/policies/groot/eagle2_hg_model/processing_eagle2_5_vl.py index 27f9b3345..7b1f67fef 100755 --- a/src/lerobot/policies/groot/eagle2_hg_model/processing_eagle2_5_vl.py +++ b/src/lerobot/policies/groot/eagle2_hg_model/processing_eagle2_5_vl.py @@ -222,6 +222,13 @@ class Eagle25VLProcessor(ProcessorMixin): videos=None, **output_kwargs["images_kwargs"], ) + if isinstance(image_inputs["pixel_values"], list): + _pv = image_inputs["pixel_values"] + if _pv and isinstance(_pv[0], list): + _pv = [t for sub in _pv for t in sub] + image_inputs["pixel_values"] = torch.stack( + [t if isinstance(t, torch.Tensor) else torch.as_tensor(t) for t in _pv] + ) num_all_tiles = image_inputs["pixel_values"].shape[0] special_placeholder = f"{self.image_start_token}{self.image_token * num_all_tiles * self.tokens_per_tile}{self.image_end_token}" unified_frame_list.append(image_inputs) @@ -233,6 +240,13 @@ class Eagle25VLProcessor(ProcessorMixin): videos=[video_list[idx_in_list]], **output_kwargs["videos_kwargs"], ) + if isinstance(video_inputs["pixel_values"], list): + _pv = video_inputs["pixel_values"] + if _pv and isinstance(_pv[0], list): + _pv = [t for sub in _pv for t in sub] + video_inputs["pixel_values"] = torch.stack( + [t if isinstance(t, torch.Tensor) else torch.as_tensor(t) for t in _pv] + ) num_all_tiles = video_inputs["pixel_values"].shape[0] image_sizes = video_inputs["image_sizes"] if timestamps_list is not None and -1 not in timestamps_list: @@ -288,8 +302,18 @@ class Eagle25VLProcessor(ProcessorMixin): text = replace_in_text(text) if len(unified_frame_list) > 0: - pixel_values = torch.cat([frame["pixel_values"] for frame in unified_frame_list]) - image_sizes = torch.cat([frame["image_sizes"] for frame in unified_frame_list]) + + def _to_tensor(v): + if isinstance(v, torch.Tensor): + return v + if isinstance(v, list): + if v and isinstance(v[0], list): + v = [t for sub in v for t in sub] + return torch.stack([t if isinstance(t, torch.Tensor) else torch.as_tensor(t) for t in v]) + return torch.as_tensor(v) + + pixel_values = torch.cat([_to_tensor(frame["pixel_values"]) for frame in unified_frame_list]) + image_sizes = torch.cat([_to_tensor(frame["image_sizes"]) for frame in unified_frame_list]) else: pixel_values = None image_sizes = None diff --git a/src/lerobot/policies/groot/groot_n1.py b/src/lerobot/policies/groot/groot_n1.py index fc753839a..abcbb8a8c 100644 --- a/src/lerobot/policies/groot/groot_n1.py +++ b/src/lerobot/policies/groot/groot_n1.py @@ -221,6 +221,7 @@ class GR00TN15(PreTrainedModel): self.action_horizon = config.action_horizon self.action_dim = config.action_dim self.compute_dtype = config.compute_dtype + self.post_init() def validate_inputs(self, inputs): # NOTE -- this should be handled internally by the model diff --git a/src/lerobot/policies/groot/modeling_groot.py b/src/lerobot/policies/groot/modeling_groot.py index 4b612bca4..2e2e9ca89 100644 --- a/src/lerobot/policies/groot/modeling_groot.py +++ b/src/lerobot/policies/groot/modeling_groot.py @@ -43,6 +43,7 @@ from torch import Tensor from lerobot.configs import FeatureType, PolicyFeature from lerobot.utils.constants import ACTION, OBS_IMAGES +from lerobot.utils.import_utils import require_package from ..pretrained import PreTrainedPolicy from .configuration_groot import GrootConfig @@ -59,6 +60,7 @@ class GrootPolicy(PreTrainedPolicy): def __init__(self, config: GrootConfig, **kwargs): """Initialize Groot policy wrapper.""" + require_package("transformers", extra="groot") super().__init__(config) config.validate_features() self.config = config diff --git a/src/lerobot/policies/multi_task_dit/modeling_multi_task_dit.py b/src/lerobot/policies/multi_task_dit/modeling_multi_task_dit.py index 8e5d1e3cb..366b271c0 100644 --- a/src/lerobot/policies/multi_task_dit/modeling_multi_task_dit.py +++ b/src/lerobot/policies/multi_task_dit/modeling_multi_task_dit.py @@ -36,7 +36,7 @@ import torch.nn.functional as F # noqa: N812 import torchvision from torch import Tensor -from lerobot.utils.import_utils import _transformers_available +from lerobot.utils.import_utils import _diffusers_available, _transformers_available, require_package from .configuration_multi_task_dit import MultiTaskDiTConfig @@ -46,6 +46,13 @@ if TYPE_CHECKING or _transformers_available: else: CLIPTextModel = None CLIPVisionModel = None + +if TYPE_CHECKING or _diffusers_available: + from diffusers.schedulers.scheduling_ddim import DDIMScheduler + from diffusers.schedulers.scheduling_ddpm import DDPMScheduler +else: + DDIMScheduler = None + DDPMScheduler = None from lerobot.utils.constants import ( ACTION, OBS_IMAGES, @@ -65,6 +72,8 @@ class MultiTaskDiTPolicy(PreTrainedPolicy): name = "multi_task_dit" def __init__(self, config: MultiTaskDiTConfig, **kwargs): + require_package("transformers", extra="multi_task_dit") + require_package("diffusers", extra="multi_task_dit") super().__init__(config) config.validate_features() self.config = config @@ -643,12 +652,6 @@ class DiffusionObjective(nn.Module): "prediction_type": config.prediction_type, } - from lerobot.utils.import_utils import require_package - - require_package("diffusers", extra="multi_task_dit") - from diffusers.schedulers.scheduling_ddim import DDIMScheduler - from diffusers.schedulers.scheduling_ddpm import DDPMScheduler - if config.noise_scheduler_type == "DDPM": self.noise_scheduler: DDPMScheduler | DDIMScheduler = DDPMScheduler(**scheduler_kwargs) elif config.noise_scheduler_type == "DDIM": diff --git a/src/lerobot/policies/pi0/modeling_pi0.py b/src/lerobot/policies/pi0/modeling_pi0.py index 22e4e6a26..3534c7ae8 100644 --- a/src/lerobot/policies/pi0/modeling_pi0.py +++ b/src/lerobot/policies/pi0/modeling_pi0.py @@ -26,7 +26,7 @@ import torch import torch.nn.functional as F # noqa: N812 from torch import Tensor, nn -from lerobot.utils.import_utils import _transformers_available +from lerobot.utils.import_utils import _transformers_available, require_package # Conditional import for type checking and lazy loading if TYPE_CHECKING or _transformers_available: @@ -947,6 +947,7 @@ class PI0Policy(PreTrainedPolicy): Args: config: Policy configuration class instance. """ + require_package("transformers", extra="pi") super().__init__(config) config.validate_features() self.config = config diff --git a/src/lerobot/policies/pi05/modeling_pi05.py b/src/lerobot/policies/pi05/modeling_pi05.py index a44817a74..56786fbcd 100644 --- a/src/lerobot/policies/pi05/modeling_pi05.py +++ b/src/lerobot/policies/pi05/modeling_pi05.py @@ -26,7 +26,7 @@ import torch import torch.nn.functional as F # noqa: N812 from torch import Tensor, nn -from lerobot.utils.import_utils import _transformers_available +from lerobot.utils.import_utils import _transformers_available, require_package # Conditional import for type checking and lazy loading if TYPE_CHECKING or _transformers_available: @@ -918,6 +918,7 @@ class PI05Policy(PreTrainedPolicy): Args: config: Policy configuration class instance. """ + require_package("transformers", extra="pi") super().__init__(config) config.validate_features() self.config = config diff --git a/src/lerobot/policies/pi0_fast/modeling_pi0_fast.py b/src/lerobot/policies/pi0_fast/modeling_pi0_fast.py index e86b8ad27..a49828ad1 100644 --- a/src/lerobot/policies/pi0_fast/modeling_pi0_fast.py +++ b/src/lerobot/policies/pi0_fast/modeling_pi0_fast.py @@ -26,7 +26,7 @@ import torch import torch.nn.functional as F # noqa: N812 from torch import Tensor, nn -from lerobot.utils.import_utils import _scipy_available, _transformers_available +from lerobot.utils.import_utils import _scipy_available, _transformers_available, require_package # Conditional import for type checking and lazy loading if TYPE_CHECKING or _scipy_available: @@ -35,7 +35,7 @@ else: idct = None if TYPE_CHECKING or _transformers_available: - from transformers import AutoTokenizer + from transformers import AutoProcessor, AutoTokenizer from transformers.models.auto import CONFIG_MAPPING from ..pi_gemma import ( @@ -44,6 +44,7 @@ if TYPE_CHECKING or _transformers_available: ) else: CONFIG_MAPPING = None + AutoProcessor = None AutoTokenizer = None PiGemmaModel = None PaliGemmaForConditionalGenerationWithPiGemma = None @@ -826,14 +827,14 @@ class PI0FastPolicy(PreTrainedPolicy): Args: config: Policy configuration class instance. """ + require_package("transformers", extra="pi") + require_package("scipy", extra="pi") super().__init__(config) config.validate_features() self.config = config # Load tokenizers first try: - from transformers import AutoProcessor, AutoTokenizer - # Load FAST tokenizer self.action_tokenizer = AutoProcessor.from_pretrained( config.action_tokenizer_name, trust_remote_code=True diff --git a/src/lerobot/policies/smolvla/modeling_smolvla.py b/src/lerobot/policies/smolvla/modeling_smolvla.py index ee3ff4db9..8ddb023da 100644 --- a/src/lerobot/policies/smolvla/modeling_smolvla.py +++ b/src/lerobot/policies/smolvla/modeling_smolvla.py @@ -62,6 +62,7 @@ from torch import Tensor, nn from lerobot.utils.constants import ACTION, OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS, OBS_STATE from lerobot.utils.device_utils import get_safe_dtype +from lerobot.utils.import_utils import require_package from ..pretrained import PreTrainedPolicy from ..rtc.modeling_rtc import RTCProcessor @@ -239,6 +240,7 @@ class SmolVLAPolicy(PreTrainedPolicy): the configuration class is used. """ + require_package("transformers", extra="smolvla") super().__init__(config) config.validate_features() self.config = config diff --git a/src/lerobot/rl/buffer.py b/src/lerobot/rl/buffer.py index 97aaa9caa..05b8419bd 100644 --- a/src/lerobot/rl/buffer.py +++ b/src/lerobot/rl/buffer.py @@ -15,6 +15,7 @@ # limitations under the License. import functools +import threading from collections.abc import Callable, Sequence from contextlib import suppress from typing import TypedDict @@ -115,6 +116,7 @@ class ReplayBuffer: self.size = 0 self.initialized = False self.optimize_memory = optimize_memory + self._lock = threading.Lock() # Track episode boundaries for memory optimization self.episode_ends = torch.zeros(capacity, dtype=torch.bool, device=storage_device) @@ -198,68 +200,75 @@ class ReplayBuffer: complementary_info: dict[str, torch.Tensor] | None = None, ): """Saves a transition, ensuring tensors are stored on the designated storage device.""" - # Initialize storage if this is the first transition - if not self.initialized: - self._initialize_storage(state=state, action=action, complementary_info=complementary_info) + with self._lock: + # Initialize storage if this is the first transition + if not self.initialized: + self._initialize_storage(state=state, action=action, complementary_info=complementary_info) - # Store the transition in pre-allocated tensors - for key in self.states: - self.states[key][self.position].copy_(state[key].squeeze(dim=0)) + # Store the transition in pre-allocated tensors + for key in self.states: + self.states[key][self.position].copy_(state[key].squeeze(dim=0)) - if not self.optimize_memory: - # Only store next_states if not optimizing memory - self.next_states[key][self.position].copy_(next_state[key].squeeze(dim=0)) + if not self.optimize_memory: + # Only store next_states if not optimizing memory + self.next_states[key][self.position].copy_(next_state[key].squeeze(dim=0)) - self.actions[self.position].copy_(action.squeeze(dim=0)) - self.rewards[self.position] = reward - self.dones[self.position] = done - self.truncateds[self.position] = truncated + self.actions[self.position].copy_(action.squeeze(dim=0)) + self.rewards[self.position] = reward + self.dones[self.position] = done + self.truncateds[self.position] = truncated - # Handle complementary_info if provided and storage is initialized - if complementary_info is not None and self.has_complementary_info: - # Store the complementary_info - for key in self.complementary_info_keys: - if key in complementary_info: - value = complementary_info[key] - if isinstance(value, torch.Tensor): - self.complementary_info[key][self.position].copy_(value.squeeze(dim=0)) - elif isinstance(value, (int | float)): - self.complementary_info[key][self.position] = value + # Handle complementary_info if provided and storage is initialized + if complementary_info is not None and self.has_complementary_info: + for key in self.complementary_info_keys: + if key in complementary_info: + value = complementary_info[key] + if isinstance(value, torch.Tensor): + self.complementary_info[key][self.position].copy_(value.squeeze(dim=0)) + elif isinstance(value, (int | float)): + self.complementary_info[key][self.position] = value - self.position = (self.position + 1) % self.capacity - self.size = min(self.size + 1, self.capacity) + self.position = (self.position + 1) % self.capacity + self.size = min(self.size + 1, self.capacity) def sample(self, batch_size: int) -> BatchTransition: """Sample a random batch of transitions and collate them into batched tensors.""" if not self.initialized: raise RuntimeError("Cannot sample from an empty buffer. Add transitions first.") - batch_size = min(batch_size, self.size) - high = max(0, self.size - 1) if self.optimize_memory and self.size < self.capacity else self.size + with self._lock: + batch_size = min(batch_size, self.size) + high = max(0, self.size - 1) if self.optimize_memory and self.size < self.capacity else self.size - # Random indices for sampling - create on the same device as storage - idx = torch.randint(low=0, high=high, size=(batch_size,), device=self.storage_device) + idx = torch.randint(low=0, high=high, size=(batch_size,), device=self.storage_device) - # Identify image keys that need augmentation - image_keys = [k for k in self.states if k.startswith(OBS_IMAGE)] if self.use_drq else [] + image_keys = [k for k in self.states if k.startswith(OBS_IMAGE)] if self.use_drq else [] - # Create batched state and next_state - batch_state = {} - batch_next_state = {} + batch_state = {} + batch_next_state = {} - # First pass: load all state tensors to target device - for key in self.states: - batch_state[key] = self.states[key][idx].to(self.device) + for key in self.states: + batch_state[key] = self.states[key][idx].to(self.device) - if not self.optimize_memory: - # Standard approach - load next_states directly - batch_next_state[key] = self.next_states[key][idx].to(self.device) - else: - # Memory-optimized approach - get next_state from the next index - next_idx = (idx + 1) % self.capacity - batch_next_state[key] = self.states[key][next_idx].to(self.device) + if not self.optimize_memory: + batch_next_state[key] = self.next_states[key][idx].to(self.device) + else: + next_idx = (idx + 1) % self.capacity + batch_next_state[key] = self.states[key][next_idx].to(self.device) + + # Sample other tensors + batch_actions = self.actions[idx].to(self.device) + batch_rewards = self.rewards[idx].to(self.device) + batch_dones = self.dones[idx].to(self.device).float() + batch_truncateds = self.truncateds[idx].to(self.device).float() + + # Sample complementary_info if available + batch_complementary_info = None + if self.has_complementary_info: + batch_complementary_info = {} + for key in self.complementary_info_keys: + batch_complementary_info[key] = self.complementary_info[key][idx].to(self.device) - # Apply image augmentation in a batched way if needed if self.use_drq and image_keys: # Concatenate all images from state and next_state all_images = [] @@ -280,19 +289,6 @@ class ReplayBuffer: # Next states start after the states at index (i*2+1)*batch_size and also take up batch_size slots batch_next_state[key] = augmented_images[(i * 2 + 1) * batch_size : (i + 1) * 2 * batch_size] - # Sample other tensors - batch_actions = self.actions[idx].to(self.device) - batch_rewards = self.rewards[idx].to(self.device) - batch_dones = self.dones[idx].to(self.device).float() - batch_truncateds = self.truncateds[idx].to(self.device).float() - - # Sample complementary_info if available - batch_complementary_info = None - if self.has_complementary_info: - batch_complementary_info = {} - for key in self.complementary_info_keys: - batch_complementary_info[key] = self.complementary_info[key][idx].to(self.device) - return BatchTransition( state=batch_state, action=batch_actions, diff --git a/src/lerobot/rl/gym_manipulator.py b/src/lerobot/rl/gym_manipulator.py index b6ff7155a..2190070f5 100644 --- a/src/lerobot/rl/gym_manipulator.py +++ b/src/lerobot/rl/gym_manipulator.py @@ -551,8 +551,8 @@ def step_env_and_process_transition( terminated = terminated or processed_action_transition[TransitionKey.DONE] truncated = truncated or processed_action_transition[TransitionKey.TRUNCATED] complementary_data = processed_action_transition[TransitionKey.COMPLEMENTARY_DATA].copy() - new_info = processed_action_transition[TransitionKey.INFO].copy() - new_info.update(info) + new_info = info.copy() + new_info.update(processed_action_transition[TransitionKey.INFO]) new_transition = create_transition( observation=obs, diff --git a/src/lerobot/robots/reachy2/robot_reachy2.py b/src/lerobot/robots/reachy2/robot_reachy2.py index ef55f71b9..ac5c9ef2f 100644 --- a/src/lerobot/robots/reachy2/robot_reachy2.py +++ b/src/lerobot/robots/reachy2/robot_reachy2.py @@ -20,7 +20,7 @@ from typing import TYPE_CHECKING, Any from lerobot.cameras import make_cameras_from_configs from lerobot.types import RobotAction, RobotObservation -from lerobot.utils.import_utils import _reachy2_sdk_available +from lerobot.utils.import_utils import _reachy2_sdk_available, require_package from ..robot import Robot from ..utils import ensure_safe_goal_position @@ -81,6 +81,7 @@ class Reachy2Robot(Robot): name = "reachy2" def __init__(self, config: Reachy2RobotConfig): + require_package("reachy2_sdk", extra="reachy2") super().__init__(config) self.config = config diff --git a/src/lerobot/robots/unitree_g1/unitree_g1.py b/src/lerobot/robots/unitree_g1/unitree_g1.py index 785861a5a..25ec32716 100644 --- a/src/lerobot/robots/unitree_g1/unitree_g1.py +++ b/src/lerobot/robots/unitree_g1/unitree_g1.py @@ -27,7 +27,7 @@ import numpy as np from lerobot.cameras import make_cameras_from_configs from lerobot.types import RobotAction, RobotObservation -from lerobot.utils.import_utils import _unitree_sdk_available +from lerobot.utils.import_utils import _unitree_sdk_available, require_package from ..robot import Robot from .config_unitree_g1 import UnitreeG1Config @@ -111,6 +111,7 @@ class UnitreeG1(Robot): name = "unitree_g1" def __init__(self, config: UnitreeG1Config): + require_package("unitree-sdk2py", extra="unitree_g1", import_name="unitree_sdk2py") super().__init__(config) logger.info("Initialize UnitreeG1...") diff --git a/src/lerobot/teleoperators/gamepad/gamepad_utils.py b/src/lerobot/teleoperators/gamepad/gamepad_utils.py index 9f94b6746..c1531ca84 100644 --- a/src/lerobot/teleoperators/gamepad/gamepad_utils.py +++ b/src/lerobot/teleoperators/gamepad/gamepad_utils.py @@ -15,9 +15,22 @@ # limitations under the License. import logging +from typing import TYPE_CHECKING + +from lerobot.utils.import_utils import _hidapi_available, _pygame_available, require_package from ..utils import TeleopEvents +if TYPE_CHECKING or _pygame_available: + import pygame +else: + pygame = None # type: ignore[assignment] + +if TYPE_CHECKING or _hidapi_available: + import hid +else: + hid = None # type: ignore[assignment] + class InputController: """Base class for input controllers that generate motion deltas.""" @@ -199,6 +212,7 @@ class GamepadController(InputController): """Generate motion deltas from gamepad input.""" def __init__(self, x_step_size=1.0, y_step_size=1.0, z_step_size=1.0, deadzone=0.1): + require_package("pygame", extra="gamepad") super().__init__(x_step_size, y_step_size, z_step_size) self.deadzone = deadzone self.joystick = None @@ -206,8 +220,6 @@ class GamepadController(InputController): def start(self): """Initialize pygame and the gamepad.""" - import pygame - pygame.init() pygame.joystick.init() @@ -230,8 +242,6 @@ class GamepadController(InputController): def stop(self): """Clean up pygame resources.""" - import pygame - if pygame.joystick.get_init(): if self.joystick: self.joystick.quit() @@ -240,8 +250,6 @@ class GamepadController(InputController): def update(self): """Process pygame events to get fresh gamepad readings.""" - import pygame - for event in pygame.event.get(): if event.type == pygame.JOYBUTTONDOWN: if event.button == 3: @@ -280,8 +288,6 @@ class GamepadController(InputController): def get_deltas(self): """Get the current movement deltas from gamepad state.""" - import pygame - try: # Read joystick axes # Left stick X and Y (typically axes 0 and 1) @@ -326,6 +332,7 @@ class GamepadControllerHID(InputController): z_scale: Scaling factor for Z-axis movement deadzone: Joystick deadzone to prevent drift """ + require_package("hidapi", extra="gamepad", import_name="hid") super().__init__(x_step_size, y_step_size, z_step_size) self.deadzone = deadzone self.device = None @@ -342,8 +349,6 @@ class GamepadControllerHID(InputController): def find_device(self): """Look for the gamepad device by vendor and product ID.""" - import hid - devices = hid.enumerate() for device in devices: device_name = device["product_string"] @@ -357,8 +362,6 @@ class GamepadControllerHID(InputController): def start(self): """Connect to the gamepad using HIDAPI.""" - import hid - self.device_info = self.find_device() if not self.device_info: self.running = False diff --git a/src/lerobot/teleoperators/homunculus/homunculus_arm.py b/src/lerobot/teleoperators/homunculus/homunculus_arm.py index 225235b59..4ceade847 100644 --- a/src/lerobot/teleoperators/homunculus/homunculus_arm.py +++ b/src/lerobot/teleoperators/homunculus/homunculus_arm.py @@ -45,7 +45,7 @@ class HomunculusArm(Teleoperator): name = "homunculus_arm" def __init__(self, config: HomunculusArmConfig): - require_package("pyserial", extra="hardware", import_name="serial") + require_package("pyserial", extra="pyserial-dep", import_name="serial") super().__init__(config) self.config = config self.serial = serial.Serial(config.port, config.baud_rate, timeout=1) diff --git a/src/lerobot/teleoperators/homunculus/homunculus_glove.py b/src/lerobot/teleoperators/homunculus/homunculus_glove.py index 655bae726..cd503c20a 100644 --- a/src/lerobot/teleoperators/homunculus/homunculus_glove.py +++ b/src/lerobot/teleoperators/homunculus/homunculus_glove.py @@ -71,7 +71,7 @@ class HomunculusGlove(Teleoperator): name = "homunculus_glove" def __init__(self, config: HomunculusGloveConfig): - require_package("pyserial", extra="hardware", import_name="serial") + require_package("pyserial", extra="pyserial-dep", import_name="serial") super().__init__(config) self.config = config self.serial = serial.Serial(config.port, config.baud_rate, timeout=1) diff --git a/src/lerobot/teleoperators/keyboard/teleop_keyboard.py b/src/lerobot/teleoperators/keyboard/teleop_keyboard.py index 0f1c7d7f1..6fc553d38 100644 --- a/src/lerobot/teleoperators/keyboard/teleop_keyboard.py +++ b/src/lerobot/teleoperators/keyboard/teleop_keyboard.py @@ -23,7 +23,7 @@ from typing import Any from lerobot.types import RobotAction from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected -from lerobot.utils.import_utils import _pynput_available +from lerobot.utils.import_utils import _pynput_available, require_package from ..teleoperator import Teleoperator from ..utils import TeleopEvents @@ -56,6 +56,7 @@ class KeyboardTeleop(Teleoperator): name = "keyboard" def __init__(self, config: KeyboardTeleopConfig): + require_package("pynput", extra="pynput-dep") super().__init__(config) self.config = config self.robot_type = config.type diff --git a/src/lerobot/teleoperators/phone/teleop_phone.py b/src/lerobot/teleoperators/phone/teleop_phone.py index f68843194..f1af248e4 100644 --- a/src/lerobot/teleoperators/phone/teleop_phone.py +++ b/src/lerobot/teleoperators/phone/teleop_phone.py @@ -21,14 +21,24 @@ import logging import threading import time +from typing import TYPE_CHECKING -import hebi import numpy as np -from teleop import Teleop from lerobot.utils.decorators import check_if_already_connected, check_if_not_connected +from lerobot.utils.import_utils import _hebi_available, _teleop_available, require_package from lerobot.utils.rotation import Rotation +if TYPE_CHECKING or _hebi_available: + import hebi +else: + hebi = None + +if TYPE_CHECKING or _teleop_available: + from teleop import Teleop +else: + Teleop = None + from ..teleoperator import Teleoperator from .config_phone import PhoneConfig, PhoneOS @@ -74,6 +84,8 @@ class IOSPhone(BasePhone, Teleoperator): name = "ios_phone" def __init__(self, config: PhoneConfig): + require_package("hebi-py", extra="phone", import_name="hebi") + require_package("teleop", extra="phone") super().__init__(config) self.config = config self._group = None @@ -213,6 +225,8 @@ class AndroidPhone(BasePhone, Teleoperator): name = "android_phone" def __init__(self, config: PhoneConfig): + require_package("hebi-py", extra="phone", import_name="hebi") + require_package("teleop", extra="phone") super().__init__(config) self.config = config self._teleop = None diff --git a/src/lerobot/teleoperators/reachy2_teleoperator/reachy2_teleoperator.py b/src/lerobot/teleoperators/reachy2_teleoperator/reachy2_teleoperator.py index db076b20f..9afb34fd7 100644 --- a/src/lerobot/teleoperators/reachy2_teleoperator/reachy2_teleoperator.py +++ b/src/lerobot/teleoperators/reachy2_teleoperator/reachy2_teleoperator.py @@ -19,7 +19,7 @@ import logging import time from typing import TYPE_CHECKING -from lerobot.utils.import_utils import _reachy2_sdk_available +from lerobot.utils.import_utils import _reachy2_sdk_available, require_package if TYPE_CHECKING or _reachy2_sdk_available: from reachy2_sdk import ReachySDK @@ -84,6 +84,7 @@ class Reachy2Teleoperator(Teleoperator): name = "reachy2_specific" def __init__(self, config: Reachy2TeleoperatorConfig): + require_package("reachy2_sdk", extra="reachy2") super().__init__(config) self.config = config diff --git a/src/lerobot/teleoperators/unitree_g1/exo_calib.py b/src/lerobot/teleoperators/unitree_g1/exo_calib.py index 05f5180ff..e977cd8b7 100644 --- a/src/lerobot/teleoperators/unitree_g1/exo_calib.py +++ b/src/lerobot/teleoperators/unitree_g1/exo_calib.py @@ -34,7 +34,7 @@ from typing import TYPE_CHECKING import numpy as np -from lerobot.utils.import_utils import _serial_available +from lerobot.utils.import_utils import _serial_available, require_package if TYPE_CHECKING or _serial_available: import serial @@ -156,6 +156,7 @@ def run_exo_calibration( """ Run interactive calibration for an exoskeleton arm. """ + require_package("pyserial", extra="unitree_g1", import_name="serial") try: import cv2 import matplotlib.pyplot as plt diff --git a/src/lerobot/teleoperators/unitree_g1/exo_serial.py b/src/lerobot/teleoperators/unitree_g1/exo_serial.py index 9b1c71891..ce5492537 100644 --- a/src/lerobot/teleoperators/unitree_g1/exo_serial.py +++ b/src/lerobot/teleoperators/unitree_g1/exo_serial.py @@ -76,7 +76,7 @@ class ExoskeletonArm: calibration: ExoskeletonCalibration | None = None def __post_init__(self): - require_package("pyserial", extra="hardware", import_name="serial") + require_package("pyserial", extra="unitree_g1", import_name="serial") if self.calibration_fpath.is_file(): self._load_calibration() diff --git a/src/lerobot/utils/import_utils.py b/src/lerobot/utils/import_utils.py index 8cd24b0fa..1ec0b6375 100644 --- a/src/lerobot/utils/import_utils.py +++ b/src/lerobot/utils/import_utils.py @@ -115,6 +115,12 @@ _feetech_sdk_available = is_package_available("feetech-servo-sdk", import_name=" _reachy2_sdk_available = is_package_available("reachy2_sdk") _can_available = is_package_available("python-can", "can") _unitree_sdk_available = is_package_available("unitree-sdk2py", "unitree_sdk2py") +_pyrealsense2_available = is_package_available("pyrealsense2") +_zmq_available = is_package_available("pyzmq", import_name="zmq") +_hebi_available = is_package_available("hebi-py", import_name="hebi") +_teleop_available = is_package_available("teleop") +_placo_available = is_package_available("placo") +_hidapi_available = is_package_available("hidapi", import_name="hid") # Data / serialization _pandas_available = is_package_available("pandas") diff --git a/tests/policies/multi_task_dit/test_multi_task_dit.py b/tests/policies/multi_task_dit/test_multi_task_dit.py index 5b70422d4..e4d456d19 100644 --- a/tests/policies/multi_task_dit/test_multi_task_dit.py +++ b/tests/policies/multi_task_dit/test_multi_task_dit.py @@ -147,6 +147,7 @@ def test_multi_task_dit_policy_forward(batch_size: int, state_dim: int, action_d ) policy = MultiTaskDiTPolicy(config=config) + policy.to(config.device) policy.train() # Use preprocessor to handle tokenization @@ -336,6 +337,7 @@ def test_multi_task_dit_policy_select_action(batch_size: int, state_dim: int, ac ) policy = MultiTaskDiTPolicy(config=config) + policy.to(config.device) policy.eval() policy.reset() # Reset queues before inference @@ -390,6 +392,7 @@ def test_multi_task_dit_policy_diffusion_objective(): config.validate_features() policy = MultiTaskDiTPolicy(config=config) + policy.to(config.device) policy.train() # Use preprocessor to handle tokenization @@ -468,6 +471,7 @@ def test_multi_task_dit_policy_flow_matching_objective(): config.validate_features() policy = MultiTaskDiTPolicy(config=config) + policy.to(config.device) policy.train() # Use preprocessor to handle tokenization @@ -533,16 +537,12 @@ def test_multi_task_dit_policy_save_and_load(tmp_path): ) policy = MultiTaskDiTPolicy(config=config) + policy.to(config.device) policy.eval() - # Get device before saving - device = next(policy.parameters()).device - policy.save_pretrained(root) loaded_policy = MultiTaskDiTPolicy.from_pretrained(root, config=config) - - # Explicitly move loaded_policy to the same device - loaded_policy.to(device) + loaded_policy.to(config.device) loaded_policy.eval() batch = create_train_batch( @@ -565,10 +565,6 @@ def test_multi_task_dit_policy_save_and_load(tmp_path): with seeded_context(12): # Process batch through preprocessor processed_batch = preprocessor(batch) - # Move batch to the same device as the policy - for key in processed_batch: - if isinstance(processed_batch[key], torch.Tensor): - processed_batch[key] = processed_batch[key].to(device) # Collect policy values before saving loss, _ = policy.forward(processed_batch) @@ -608,6 +604,7 @@ def test_multi_task_dit_policy_get_optim_params(): ) policy = MultiTaskDiTPolicy(config=config) + policy.to(config.device) param_groups = policy.get_optim_params() # Should have 2 parameter groups: non-vision and vision encoder diff --git a/tests/teleoperators/test_reachy2_teleoperator.py b/tests/teleoperators/test_reachy2_teleoperator.py index dd8c5904c..f0274cb75 100644 --- a/tests/teleoperators/test_reachy2_teleoperator.py +++ b/tests/teleoperators/test_reachy2_teleoperator.py @@ -18,6 +18,11 @@ from unittest.mock import MagicMock, patch import pytest +from lerobot.utils.import_utils import is_package_available + +if not is_package_available("reachy2_sdk"): + pytest.skip("reachy2_sdk not available", allow_module_level=True) + from lerobot.teleoperators.reachy2_teleoperator import ( REACHY2_ANTENNAS_JOINTS, REACHY2_L_ARM_JOINTS,