mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 10:16:09 +00:00
refactor(pipeline): Simplify observation and padding data handling in batch transitions
This commit is contained in:
@@ -183,15 +183,13 @@ def _default_batch_to_transition(batch: dict[str, Any]) -> EnvTransition: # noq
|
|||||||
metadata without breaking the processor.
|
metadata without breaking the processor.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Handle observation and observation.* keys
|
# Extract observation keys
|
||||||
observation_keys = {k: v for k, v in batch.items() if k.startswith("observation.")}
|
observation_keys = {k: v for k, v in batch.items() if k.startswith("observation.")}
|
||||||
|
observation = observation_keys if observation_keys else None
|
||||||
|
|
||||||
observation = None
|
# Extract padding keys for complementary data
|
||||||
if observation_keys:
|
pad_keys = {k: v for k, v in batch.items() if "_is_pad" in k}
|
||||||
observation = {}
|
complementary_data = pad_keys if pad_keys else {}
|
||||||
# Add observation.* keys to the observation dict
|
|
||||||
for key, value in observation_keys.items():
|
|
||||||
observation[key] = value
|
|
||||||
|
|
||||||
return (
|
return (
|
||||||
observation,
|
observation,
|
||||||
@@ -200,7 +198,7 @@ def _default_batch_to_transition(batch: dict[str, Any]) -> EnvTransition: # noq
|
|||||||
batch.get("next.done", False),
|
batch.get("next.done", False),
|
||||||
batch.get("next.truncated", False),
|
batch.get("next.truncated", False),
|
||||||
batch.get("info", {}),
|
batch.get("info", {}),
|
||||||
{},
|
complementary_data,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -216,7 +214,7 @@ def _default_transition_to_batch(transition: EnvTransition) -> dict[str, Any]:
|
|||||||
done,
|
done,
|
||||||
truncated,
|
truncated,
|
||||||
info,
|
info,
|
||||||
_,
|
complementary_data,
|
||||||
) = transition
|
) = transition
|
||||||
|
|
||||||
batch = {
|
batch = {
|
||||||
@@ -227,11 +225,15 @@ def _default_transition_to_batch(transition: EnvTransition) -> dict[str, Any]:
|
|||||||
"info": info,
|
"info": info,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# Add padding data from complementary_data
|
||||||
|
if complementary_data:
|
||||||
|
pad_data = {k: v for k, v in complementary_data.items() if "_is_pad" in k}
|
||||||
|
batch.update(pad_data)
|
||||||
|
|
||||||
# Handle observation - flatten dict to observation.* keys if it's a dict
|
# Handle observation - flatten dict to observation.* keys if it's a dict
|
||||||
if isinstance(observation, dict):
|
if isinstance(observation, dict):
|
||||||
# Check if this looks like a dict that was created from observation.* keys
|
batch.update(observation)
|
||||||
for key, value in observation.items():
|
|
||||||
batch[key] = value
|
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user