mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-25 10:46:01 +00:00
fix: populate config from dataset metadata
This commit is contained in:
@@ -20,8 +20,10 @@ from typing import Any
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from lerobot.configs import FeatureType
|
||||||
from lerobot.configs.rewards import RewardModelConfig
|
from lerobot.configs.rewards import RewardModelConfig
|
||||||
from lerobot.processor import PolicyAction, PolicyProcessorPipeline
|
from lerobot.processor import PolicyAction, PolicyProcessorPipeline
|
||||||
|
from lerobot.utils.feature_utils import dataset_to_policy_features
|
||||||
|
|
||||||
from .classifier.configuration_classifier import RewardClassifierConfig
|
from .classifier.configuration_classifier import RewardClassifierConfig
|
||||||
from .distributional_value_function.configuration_distributional_value_function import DistributionalVFConfig
|
from .distributional_value_function.configuration_distributional_value_function import DistributionalVFConfig
|
||||||
@@ -147,6 +149,13 @@ def make_reward_model(cfg: RewardModelConfig, **kwargs) -> PreTrainedRewardModel
|
|||||||
Returns:
|
Returns:
|
||||||
An instantiated and device-placed reward model.
|
An instantiated and device-placed reward model.
|
||||||
"""
|
"""
|
||||||
|
dataset_meta = kwargs.get("dataset_meta")
|
||||||
|
if dataset_meta is not None and not cfg.input_features:
|
||||||
|
features = dataset_to_policy_features(dataset_meta.features)
|
||||||
|
cfg.input_features = {
|
||||||
|
key: feature for key, feature in features.items() if feature.type is not FeatureType.ACTION
|
||||||
|
}
|
||||||
|
|
||||||
reward_cls = get_reward_model_class(cfg.type)
|
reward_cls = get_reward_model_class(cfg.type)
|
||||||
|
|
||||||
kwargs["config"] = cfg
|
kwargs["config"] = cfg
|
||||||
|
|||||||
@@ -0,0 +1,46 @@
|
|||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from lerobot.configs import FeatureType
|
||||||
|
from lerobot.rewards.factory import make_reward_model
|
||||||
|
from lerobot.rewards.temporal_siglip_value_function.configuration_temporal_siglip_value_function import (
|
||||||
|
TemporalSiglipVFConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_reward_factory_populates_input_features_from_dataset_meta(monkeypatch):
|
||||||
|
from lerobot.rewards import factory
|
||||||
|
|
||||||
|
class FakeReward(nn.Module):
|
||||||
|
def __init__(self, config, **kwargs):
|
||||||
|
super().__init__()
|
||||||
|
self.config = config
|
||||||
|
|
||||||
|
monkeypatch.setattr(factory, "get_reward_model_class", lambda name: FakeReward)
|
||||||
|
metadata = SimpleNamespace(
|
||||||
|
features={
|
||||||
|
"observation.images.top": {
|
||||||
|
"dtype": "video",
|
||||||
|
"shape": [480, 640, 3],
|
||||||
|
"names": ["height", "width", "channel"],
|
||||||
|
},
|
||||||
|
"observation.state": {
|
||||||
|
"dtype": "float32",
|
||||||
|
"shape": [14],
|
||||||
|
"names": [f"joint_{index}" for index in range(14)],
|
||||||
|
},
|
||||||
|
"action": {
|
||||||
|
"dtype": "float32",
|
||||||
|
"shape": [14],
|
||||||
|
"names": [f"joint_{index}" for index in range(14)],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
config = TemporalSiglipVFConfig(device="cpu")
|
||||||
|
model = make_reward_model(config, dataset_meta=metadata)
|
||||||
|
|
||||||
|
assert model.config.input_features["observation.images.top"].type is FeatureType.VISUAL
|
||||||
|
assert model.config.input_features["observation.images.top"].shape == (3, 480, 640)
|
||||||
|
assert model.config.input_features["observation.state"].type is FeatureType.STATE
|
||||||
|
assert "action" not in model.config.input_features
|
||||||
Reference in New Issue
Block a user