mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-27 11:46:04 +00:00
fix(processor): rename temporal padding metadata
This commit is contained in:
@@ -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}
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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={})
|
||||
|
||||
Reference in New Issue
Block a user