mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 11:16:00 +00:00
refactor(rewards): enhance RewardModelConfig with device handling and delta indices properties
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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}"
|
||||||
|
|||||||
Reference in New Issue
Block a user