fix: populate config from dataset metadata

This commit is contained in:
Khalil Meftah
2026-07-23 22:17:39 +02:00
parent d015474483
commit 39284f43c7
2 changed files with 55 additions and 0 deletions
+9
View File
@@ -20,8 +20,10 @@ from typing import Any
import torch
from lerobot.configs import FeatureType
from lerobot.configs.rewards import RewardModelConfig
from lerobot.processor import PolicyAction, PolicyProcessorPipeline
from lerobot.utils.feature_utils import dataset_to_policy_features
from .classifier.configuration_classifier import RewardClassifierConfig
from .distributional_value_function.configuration_distributional_value_function import DistributionalVFConfig
@@ -147,6 +149,13 @@ def make_reward_model(cfg: RewardModelConfig, **kwargs) -> PreTrainedRewardModel
Returns:
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)
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