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
|
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
|
@dataclass
|
||||||
@ProcessorStepRegistry.register(name="rename_observations_processor")
|
@ProcessorStepRegistry.register(name="rename_observations_processor")
|
||||||
@@ -41,14 +58,7 @@ class RenameObservationsProcessorStep(ObservationProcessorStep):
|
|||||||
rename_map: dict[str, str] = field(default_factory=dict)
|
rename_map: dict[str, str] = field(default_factory=dict)
|
||||||
|
|
||||||
def observation(self, observation):
|
def observation(self, observation):
|
||||||
processed_obs = {}
|
return rename_transition_keys(observation, self.rename_map)
|
||||||
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
|
|
||||||
|
|
||||||
def get_config(self) -> dict[str, Any]:
|
def get_config(self) -> dict[str, Any]:
|
||||||
return {"rename_map": self.rename_map}
|
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.jobs import submit_to_hf
|
||||||
from lerobot.optim.factory import make_optimizer_and_scheduler
|
from lerobot.optim.factory import make_optimizer_and_scheduler
|
||||||
from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_processors
|
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.rewards import make_reward_pre_post_processors
|
||||||
from lerobot.utils.collate import lerobot_collate_fn
|
from lerobot.utils.collate import lerobot_collate_fn
|
||||||
from lerobot.utils.import_utils import register_third_party_plugins
|
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:
|
if cam_key in batch and batch[cam_key].dtype == torch.uint8:
|
||||||
batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0
|
batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0
|
||||||
if cfg.rename_map:
|
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)
|
batch = preprocessor(batch)
|
||||||
train_tracker.dataloading_s = time.perf_counter() - start_time
|
train_tracker.dataloading_s = time.perf_counter() - start_time
|
||||||
|
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ from lerobot.processor import (
|
|||||||
TransitionKey,
|
TransitionKey,
|
||||||
)
|
)
|
||||||
from lerobot.processor.converters import create_transition, identity_transition
|
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 lerobot.utils.constants import ACTION, OBS_IMAGE, OBS_IMAGES, OBS_STATE
|
||||||
from tests.conftest import assert_contract_is_typed
|
from tests.conftest import assert_contract_is_typed
|
||||||
|
|
||||||
@@ -64,6 +64,22 @@ def test_basic_renaming():
|
|||||||
assert processed_obs["unchanged_key"] == "keep_me"
|
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():
|
def test_empty_rename_map():
|
||||||
"""Test processor with empty rename map (should pass through unchanged)."""
|
"""Test processor with empty rename map (should pass through unchanged)."""
|
||||||
processor = RenameObservationsProcessorStep(rename_map={})
|
processor = RenameObservationsProcessorStep(rename_map={})
|
||||||
|
|||||||
Reference in New Issue
Block a user