refactor(rewards): enhance RewardModelConfig with device handling and delta indices properties

This commit is contained in:
Khalil Meftah
2026-04-23 14:05:34 +02:00
parent 9f739ac1ff
commit 0be513c250
3 changed files with 27 additions and 4 deletions
+19
View File
@@ -30,6 +30,7 @@ from huggingface_hub.errors import HfHubHTTPError
from lerobot.configs.types import PolicyFeature from lerobot.configs.types import PolicyFeature
from lerobot.optim.optimizers import OptimizerConfig from lerobot.optim.optimizers import OptimizerConfig
from lerobot.optim.schedulers import LRSchedulerConfig 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 from lerobot.utils.hub import HubMixin
T = TypeVar("T", bound="RewardModelConfig") T = TypeVar("T", bound="RewardModelConfig")
@@ -63,6 +64,12 @@ class RewardModelConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC):
tags: list[str] | None = None tags: list[str] | None = None
private: bool | 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 @property
def type(self) -> str: def type(self) -> str:
choice_name = self.get_choice_name(self.__class__) 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)}") raise TypeError(f"Expected string from get_choice_name, got {type(choice_name)}")
return 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 @abc.abstractmethod
def get_optimizer_preset(self) -> OptimizerConfig: def get_optimizer_preset(self) -> OptimizerConfig:
raise NotImplementedError raise NotImplementedError
+7 -4
View File
@@ -19,6 +19,7 @@ from pprint import pformat
import torch import torch
from lerobot.configs import PreTrainedConfig from lerobot.configs import PreTrainedConfig
from lerobot.configs.rewards import RewardModelConfig
from lerobot.configs.train import TrainPipelineConfig from lerobot.configs.train import TrainPipelineConfig
from lerobot.transforms import ImageTransforms from lerobot.transforms import ImageTransforms
from lerobot.utils.constants import ACTION, IMAGENET_STATS, OBS_PREFIX, REWARD from lerobot.utils.constants import ACTION, IMAGENET_STATS, OBS_PREFIX, REWARD
@@ -30,12 +31,14 @@ from .streaming_dataset import StreamingLeRobotDataset
def resolve_delta_timestamps( def resolve_delta_timestamps(
cfg: PreTrainedConfig, ds_meta: LeRobotDatasetMetadata cfg: PreTrainedConfig | RewardModelConfig, ds_meta: LeRobotDatasetMetadata
) -> dict[str, list] | None: ) -> 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: 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 ds_meta (LeRobotDatasetMetadata): The dataset from which features and fps are used to build
delta_timestamps against. delta_timestamps against.
@@ -82,7 +85,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
ds_meta = LeRobotDatasetMetadata( ds_meta = LeRobotDatasetMetadata(
cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision 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: if not cfg.dataset.streaming:
dataset = LeRobotDataset( dataset = LeRobotDataset(
cfg.dataset.repo_id, cfg.dataset.repo_id,
@@ -108,6 +108,7 @@ class SARMConfig(RewardModelConfig):
) )
def __post_init__(self): def __post_init__(self):
super().__post_init__()
if self.annotation_mode not in ["single_stage", "dense_only", "dual"]: if self.annotation_mode not in ["single_stage", "dense_only", "dual"]:
raise ValueError( raise ValueError(
f"annotation_mode must be 'single_stage', 'dense_only', or 'dual', got {self.annotation_mode}" f"annotation_mode must be 'single_stage', 'dense_only', or 'dual', got {self.annotation_mode}"