diff --git a/src/lerobot/configs/rewards.py b/src/lerobot/configs/rewards.py index 89fc51574..d495160bf 100644 --- a/src/lerobot/configs/rewards.py +++ b/src/lerobot/configs/rewards.py @@ -30,6 +30,7 @@ from huggingface_hub.errors import HfHubHTTPError from lerobot.configs.types import PolicyFeature from lerobot.optim.optimizers import OptimizerConfig from lerobot.optim.schedulers import LRSchedulerConfig +from lerobot.utils.device_utils import auto_select_torch_device, is_torch_device_available from lerobot.utils.hub import HubMixin T = TypeVar("T", bound="RewardModelConfig") @@ -63,6 +64,12 @@ class RewardModelConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC): tags: list[str] | None = None private: bool | None = None + def __post_init__(self) -> None: + if not self.device or not is_torch_device_available(self.device): + auto_device = auto_select_torch_device() + logger.warning(f"Device '{self.device}' is not available. Switching to '{auto_device}'.") + self.device = auto_device.type + @property def type(self) -> str: choice_name = self.get_choice_name(self.__class__) @@ -70,6 +77,18 @@ class RewardModelConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC): raise TypeError(f"Expected string from get_choice_name, got {type(choice_name)}") return choice_name + @property + def observation_delta_indices(self) -> list | None: # type: ignore[type-arg] + return None + + @property + def action_delta_indices(self) -> list | None: # type: ignore[type-arg] + return None + + @property + def reward_delta_indices(self) -> list | None: # type: ignore[type-arg] + return None + @abc.abstractmethod def get_optimizer_preset(self) -> OptimizerConfig: raise NotImplementedError diff --git a/src/lerobot/datasets/factory.py b/src/lerobot/datasets/factory.py index 040cba5cb..9ae57b0d7 100644 --- a/src/lerobot/datasets/factory.py +++ b/src/lerobot/datasets/factory.py @@ -19,6 +19,7 @@ from pprint import pformat import torch from lerobot.configs import PreTrainedConfig +from lerobot.configs.rewards import RewardModelConfig from lerobot.configs.train import TrainPipelineConfig from lerobot.transforms import ImageTransforms from lerobot.utils.constants import ACTION, IMAGENET_STATS, OBS_PREFIX, REWARD @@ -30,12 +31,14 @@ from .streaming_dataset import StreamingLeRobotDataset def resolve_delta_timestamps( - cfg: PreTrainedConfig, ds_meta: LeRobotDatasetMetadata + cfg: PreTrainedConfig | RewardModelConfig, ds_meta: LeRobotDatasetMetadata ) -> dict[str, list] | None: - """Resolves delta_timestamps by reading from the 'delta_indices' properties of the PreTrainedConfig. + """Resolves delta_timestamps by reading from the 'delta_indices' properties of the config. Args: - cfg (PreTrainedConfig): The PreTrainedConfig to read delta_indices from. + cfg (PreTrainedConfig | RewardModelConfig): The config to read delta_indices from. Both + ``PreTrainedConfig`` and concrete ``RewardModelConfig`` subclasses expose the + ``{observation,action,reward}_delta_indices`` properties used below. ds_meta (LeRobotDatasetMetadata): The dataset from which features and fps are used to build delta_timestamps against. @@ -82,7 +85,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas ds_meta = LeRobotDatasetMetadata( cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision ) - delta_timestamps = resolve_delta_timestamps(cfg.policy, ds_meta) + delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta) if not cfg.dataset.streaming: dataset = LeRobotDataset( cfg.dataset.repo_id, diff --git a/src/lerobot/rewards/sarm/configuration_sarm.py b/src/lerobot/rewards/sarm/configuration_sarm.py index b06979c8e..0d1f727f7 100644 --- a/src/lerobot/rewards/sarm/configuration_sarm.py +++ b/src/lerobot/rewards/sarm/configuration_sarm.py @@ -108,6 +108,7 @@ class SARMConfig(RewardModelConfig): ) def __post_init__(self): + super().__post_init__() if self.annotation_mode not in ["single_stage", "dense_only", "dual"]: raise ValueError( f"annotation_mode must be 'single_stage', 'dense_only', or 'dual', got {self.annotation_mode}"