diff --git a/src/lerobot/processor/rename_processor.py b/src/lerobot/processor/rename_processor.py index 5ffec6868..c4385954f 100644 --- a/src/lerobot/processor/rename_processor.py +++ b/src/lerobot/processor/rename_processor.py @@ -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} diff --git a/src/lerobot/scripts/lerobot_train.py b/src/lerobot/scripts/lerobot_train.py index 5fbe176fb..0859a7fd8 100644 --- a/src/lerobot/scripts/lerobot_train.py +++ b/src/lerobot/scripts/lerobot_train.py @@ -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 diff --git a/tests/processor/test_rename_processor.py b/tests/processor/test_rename_processor.py index efb9f9328..ee52b5139 100644 --- a/tests/processor/test_rename_processor.py +++ b/tests/processor/test_rename_processor.py @@ -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={})