| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| import functools |
| import traceback |
|
|
| import draccus.wrappers.docstring as _draccus_docstring |
| import pytest |
|
|
| from lerobot.configs.types import FeatureType, PipelineFeatureType, PolicyFeature |
| from lerobot.utils.import_utils import is_package_available |
| from tests.utils import DEVICE |
|
|
| |
| |
| |
| |
| |
| _draccus_docstring.get_attribute_docstring = functools.cache(_draccus_docstring.get_attribute_docstring) |
|
|
| |
| |
| |
| pytest_plugins = [ |
| "tests.fixtures.optimizers", |
| ] |
|
|
| if is_package_available("datasets"): |
| pytest_plugins += [ |
| "tests.fixtures.dataset_factories", |
| "tests.fixtures.files", |
| "tests.fixtures.hub", |
| ] |
|
|
|
|
| def pytest_collection_finish(): |
| print(f"\nTesting with {DEVICE=}") |
|
|
|
|
| def _is_serial_exception(exc: Exception) -> bool: |
| """Check if an exception is a SerialException without requiring pyserial.""" |
| if not is_package_available("pyserial", import_name="serial"): |
| return False |
| from serial import SerialException |
|
|
| return isinstance(exc, SerialException) |
|
|
|
|
| def _check_component_availability(component_type, available_components, make_component): |
| """Generic helper to check if a hardware component is available""" |
| if component_type not in available_components: |
| raise ValueError( |
| f"The {component_type} type is not valid. Expected one of these '{available_components}'" |
| ) |
|
|
| try: |
| component = make_component(component_type) |
| component.connect() |
| del component |
| return True |
|
|
| except Exception as e: |
| print(f"\nA {component_type} is not available.") |
|
|
| if isinstance(e, ModuleNotFoundError): |
| print(f"\nInstall module '{e.name}'") |
| elif _is_serial_exception(e): |
| print("\nNo physical device detected.") |
| elif isinstance(e, ValueError) and "camera_index" in str(e): |
| print("\nNo physical camera detected.") |
| else: |
| traceback.print_exc() |
|
|
| return False |
|
|
|
|
| @pytest.fixture |
| def patch_builtins_input(monkeypatch): |
| def print_text(text=None): |
| if text is not None: |
| print(text) |
|
|
| monkeypatch.setattr("builtins.input", print_text) |
|
|
|
|
| @pytest.fixture |
| def policy_feature_factory(): |
| """PolicyFeature factory""" |
|
|
| def _pf(ft: FeatureType, shape: tuple[int, ...]) -> PolicyFeature: |
| return PolicyFeature(type=ft, shape=shape) |
|
|
| return _pf |
|
|
|
|
| def assert_contract_is_typed(features: dict[PipelineFeatureType, dict[str, PolicyFeature]]) -> None: |
| assert isinstance(features, dict) |
| assert all(isinstance(k, PipelineFeatureType) for k in features) |
| assert all(isinstance(v, dict) for v in features.values()) |
| assert all(all(isinstance(nk, str) for nk in v) for v in features.values()) |
| assert all(all(isinstance(nv, PolicyFeature) for nv in v.values()) for v in features.values()) |
|
|