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 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