fix(processor): rename temporal padding metadata

This commit is contained in:
Pepijn
2026-07-15 20:05:25 +02:00
parent 899434903e
commit 994ef22b22
3 changed files with 37 additions and 10 deletions
+18 -8
View File
@@ -21,6 +21,23 @@ from lerobot.configs import PipelineFeatureType, PolicyFeature
from .pipeline import ObservationProcessorStep, ProcessorStepRegistry
_AUXILIARY_KEY_SUFFIXES = ("_is_pad", "_padding_mask")
def rename_transition_key(key: str, rename_map: dict[str, str]) -> str:
"""Rename a feature key and any sampling metadata derived from it."""
if key in rename_map:
return rename_map[key]
for suffix in _AUXILIARY_KEY_SUFFIXES:
if key.endswith(suffix) and key[: -len(suffix)] in rename_map:
return f"{rename_map[key[: -len(suffix)]]}{suffix}"
return key
def rename_transition_keys(data: dict[str, Any], rename_map: dict[str, str]) -> dict[str, Any]:
"""Rename all transition keys, including delta-sampling padding masks."""
return {rename_transition_key(key, rename_map): value for key, value in data.items()}
@dataclass
@ProcessorStepRegistry.register(name="rename_observations_processor")
@@ -41,14 +58,7 @@ class RenameObservationsProcessorStep(ObservationProcessorStep):
rename_map: dict[str, str] = field(default_factory=dict)
def observation(self, observation):
processed_obs = {}
for key, value in observation.items():
if key in self.rename_map:
processed_obs[self.rename_map[key]] = value
else:
processed_obs[key] = value
return processed_obs
return rename_transition_keys(observation, self.rename_map)
def get_config(self) -> dict[str, Any]:
return {"rename_map": self.rename_map}
+2 -1
View File
@@ -57,6 +57,7 @@ from lerobot.envs import close_envs, make_env, make_env_pre_post_processors
from lerobot.jobs import submit_to_hf
from lerobot.optim.factory import make_optimizer_and_scheduler
from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_processors
from lerobot.processor.rename_processor import rename_transition_keys
from lerobot.rewards import make_reward_pre_post_processors
from lerobot.utils.collate import lerobot_collate_fn
from lerobot.utils.import_utils import register_third_party_plugins
@@ -610,7 +611,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
if cam_key in batch and batch[cam_key].dtype == torch.uint8:
batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0
if cfg.rename_map:
batch = {cfg.rename_map.get(key, key): value for key, value in batch.items()}
batch = rename_transition_keys(batch, cfg.rename_map)
batch = preprocessor(batch)
train_tracker.dataloading_s = time.perf_counter() - start_time
+17 -1
View File
@@ -27,7 +27,7 @@ from lerobot.processor import (
TransitionKey,
)
from lerobot.processor.converters import create_transition, identity_transition
from lerobot.processor.rename_processor import rename_stats
from lerobot.processor.rename_processor import rename_stats, rename_transition_keys
from lerobot.utils.constants import ACTION, OBS_IMAGE, OBS_IMAGES, OBS_STATE
from tests.conftest import assert_contract_is_typed
@@ -64,6 +64,22 @@ def test_basic_renaming():
assert processed_obs["unchanged_key"] == "keep_me"
def test_renaming_preserves_feature_suffixes_for_sampling_metadata():
data = {
"image": torch.zeros(1),
"image_is_pad": torch.ones(1, dtype=torch.bool),
"image_padding_mask": torch.ones(1, dtype=torch.bool),
}
result = rename_transition_keys(data, {"image": "observation.images.camera1"})
assert set(result) == {
"observation.images.camera1",
"observation.images.camera1_is_pad",
"observation.images.camera1_padding_mask",
}
def test_empty_rename_map():
"""Test processor with empty rename map (should pass through unchanged)."""
processor = RenameObservationsProcessorStep(rename_map={})