docs(processor): bring src/lerobot/processor/ to 100% docstring coverage

Documents every remaining public class/function across the 15 files in
lerobot.processor (pipeline, batch/device/normalize/rename/observation
steps, env-specific and HIL steps, action bridges, tokenization,
language, converters, factory functions, and the normalization
migration script), reformats pre-existing docstrings to the
machine-checked Args:/Returns:/Raises: standard, and expands
docs/source/api/processor.mdx from 3 documented classes to the full
public surface.

Adds lerobot.processor to check_docstrings.py's MODULES_TO_CHECK
ratchet and removes the module's ruff D-ignore. Docstrings and doc
comment reformatting only; no behavioral changes.
This commit is contained in:
CarolinePascal
2026-08-07 11:20:01 +02:00
parent 741005d719
commit 06364ce436
21 changed files with 573 additions and 306 deletions
+240 -5
View File
@@ -7,14 +7,249 @@ See [Introduction to Robot Processors](../introduction_processors) for the conce
[Implement your own processor](../implement_your_own_processor) to write a step, and [Implement your own processor](../implement_your_own_processor) to write a step, and
[Debug your processor pipeline](../debug_processor_pipeline) when a pipeline misbehaves. [Debug your processor pipeline](../debug_processor_pipeline) when a pipeline misbehaves.
## ProcessorStep ## Core pipeline
[[autodoc]] lerobot.processor.pipeline.ProcessorStep [[autodoc]] lerobot.processor.pipeline.ProcessorStep
## DataProcessorPipeline
[[autodoc]] lerobot.processor.pipeline.DataProcessorPipeline [[autodoc]] lerobot.processor.pipeline.DataProcessorPipeline
## PolicyProcessorPipeline
[[autodoc]] lerobot.processor.pipeline.PolicyProcessorPipeline [[autodoc]] lerobot.processor.pipeline.PolicyProcessorPipeline
`RobotProcessorPipeline` is a type alias for `DataProcessorPipeline[TInput, TOutput]`, used for the
teleop-action, robot-action and robot-observation pipelines that don't go through a policy.
[[autodoc]] lerobot.processor.pipeline.ProcessorKwargs
[[autodoc]] lerobot.processor.pipeline.ProcessorStepRegistry
[[autodoc]] lerobot.processor.pipeline.ProcessorMigrationError
### Typed base steps
Each subclasses `ProcessorStep` to implement one part of an `EnvTransition` (observation, action, reward,
etc.), leaving the rest of the transition untouched by default.
[[autodoc]] lerobot.processor.pipeline.ObservationProcessorStep
[[autodoc]] lerobot.processor.pipeline.ActionProcessorStep
[[autodoc]] lerobot.processor.pipeline.RobotActionProcessorStep
[[autodoc]] lerobot.processor.pipeline.PolicyActionProcessorStep
[[autodoc]] lerobot.processor.pipeline.RewardProcessorStep
[[autodoc]] lerobot.processor.pipeline.DoneProcessorStep
[[autodoc]] lerobot.processor.pipeline.TruncatedProcessorStep
[[autodoc]] lerobot.processor.pipeline.InfoProcessorStep
[[autodoc]] lerobot.processor.pipeline.ComplementaryDataProcessorStep
[[autodoc]] lerobot.processor.pipeline.IdentityProcessorStep
## Factory functions
Build the canonical processor pipelines used by policies and robots.
[[autodoc]] lerobot.processor.factory.make_default_processors
[[autodoc]] lerobot.processor.factory.make_default_teleop_action_processor
[[autodoc]] lerobot.processor.factory.make_default_robot_action_processor
[[autodoc]] lerobot.processor.factory.make_default_robot_observation_processor
[[autodoc]] lerobot.processor.factory.DefaultPolicyProcessorSteps
[[autodoc]] lerobot.processor.factory.make_default_policy_processor_steps
[[autodoc]] lerobot.processor.factory.make_policy_processor_pipelines
[[autodoc]] lerobot.processor.factory.make_default_pre_post_processors
## Batch dimension
[[autodoc]] lerobot.processor.batch_processor.AddBatchDimensionProcessorStep
- all
[[autodoc]] lerobot.processor.batch_processor.AddBatchDimensionObservationStep
[[autodoc]] lerobot.processor.batch_processor.AddBatchDimensionActionStep
[[autodoc]] lerobot.processor.batch_processor.AddBatchDimensionComplementaryDataStep
## Device
[[autodoc]] lerobot.processor.device_processor.DeviceProcessorStep
- all
## Normalization
[[autodoc]] lerobot.processor.normalize_processor.NormalizerProcessorStep
- all
[[autodoc]] lerobot.processor.normalize_processor.UnnormalizerProcessorStep
- all
[[autodoc]] lerobot.processor.normalize_processor.hotswap_stats
## Renaming
[[autodoc]] lerobot.processor.rename_processor.RenameObservationsProcessorStep
- all
[[autodoc]] lerobot.processor.rename_processor.rename_stats
## Vanilla observation processing
Converts standard Gymnasium observations (`pixels`, `agent_pos`, `environment_state`) to the LeRobot format.
[[autodoc]] lerobot.processor.observation_processor.VanillaObservationProcessorStep
- all
## Environment-specific observation processing
[[autodoc]] lerobot.processor.env_processor.LiberoProcessorStep
- all
[[autodoc]] lerobot.processor.env_processor.IsaaclabArenaProcessorStep
- all
## NumPy / PyTorch action conversion
[[autodoc]] lerobot.processor.gym_action_processor.Torch2NumpyActionProcessorStep
- all
[[autodoc]] lerobot.processor.gym_action_processor.Numpy2TorchActionProcessorStep
- all
## Delta and relative actions
[[autodoc]] lerobot.processor.delta_action_processor.MapTensorToDeltaActionDictStep
- all
[[autodoc]] lerobot.processor.delta_action_processor.MapDeltaActionToRobotActionStep
- all
[[autodoc]] lerobot.processor.relative_action_processor.RelativeActionsProcessorStep
- all
[[autodoc]] lerobot.processor.relative_action_processor.AbsoluteActionsProcessorStep
- all
[[autodoc]] lerobot.processor.relative_action_processor.to_relative_actions
[[autodoc]] lerobot.processor.relative_action_processor.to_absolute_actions
## Policy/robot action bridge
Converts between a robot's per-motor action dict and a policy's stacked action tensor.
[[autodoc]] lerobot.processor.policy_robot_bridge.RobotActionToPolicyActionProcessorStep
- all
[[autodoc]] lerobot.processor.policy_robot_bridge.PolicyActionToRobotActionProcessorStep
- all
## Tokenization
[[autodoc]] lerobot.processor.tokenizer_processor.TokenizerProcessorStep
- all
[[autodoc]] lerobot.processor.tokenizer_processor.ActionTokenizerProcessorStep
- all
## Language
[[autodoc]] lerobot.processor.newline_task_processor.NewLineTaskProcessorStep
- all
`RenderMessagesStep` requires the `[dataset]` extra and is not re-exported from `lerobot.processor`; import
it directly from `lerobot.processor.render_messages_processor`.
[[autodoc]] lerobot.processor.render_messages_processor.RenderMessagesStep
- all
## Human-in-the-loop (HIL)
Steps supporting human-in-the-loop RL: teleop event/action bookkeeping, time limits, image preprocessing,
the `gym-hil` adapter, gripper penalties, intervention handling, and a learned reward classifier.
[[autodoc]] lerobot.processor.hil_processor.HasTeleopEvents
[[autodoc]] lerobot.processor.hil_processor.AddTeleopActionAsComplimentaryDataStep
- all
[[autodoc]] lerobot.processor.hil_processor.AddTeleopEventsAsInfoStep
- all
[[autodoc]] lerobot.processor.hil_processor.ImageCropResizeProcessorStep
- all
[[autodoc]] lerobot.processor.hil_processor.TimeLimitProcessorStep
- all
[[autodoc]] lerobot.processor.hil_processor.GymHILAdapterProcessorStep
- all
[[autodoc]] lerobot.processor.hil_processor.GripperPenaltyProcessorStep
- all
[[autodoc]] lerobot.processor.hil_processor.InterventionActionProcessorStep
- all
[[autodoc]] lerobot.processor.hil_processor.RewardClassifierProcessorStep
- all
## Transition converters
Convert between an `EnvTransition` and the raw dict formats used by robots, policies, and dataset batches.
[[autodoc]] lerobot.processor.converters.create_transition
[[autodoc]] lerobot.processor.converters.identity_transition
[[autodoc]] lerobot.processor.converters.to_tensor
[[autodoc]] lerobot.processor.converters.from_tensor_to_numpy
[[autodoc]] lerobot.processor.converters.robot_action_to_transition
[[autodoc]] lerobot.processor.converters.robot_action_observation_to_transition
[[autodoc]] lerobot.processor.converters.observation_to_transition
[[autodoc]] lerobot.processor.converters.policy_action_to_transition
[[autodoc]] lerobot.processor.converters.batch_to_transition
[[autodoc]] lerobot.processor.converters.transition_to_robot_action
[[autodoc]] lerobot.processor.converters.transition_to_policy_action
[[autodoc]] lerobot.processor.converters.transition_to_observation
[[autodoc]] lerobot.processor.converters.transition_to_batch
## Migrating legacy policies
Standalone script to migrate a pretrained policy with built-in normalization layers to the processor
pipeline system. See the module docstring for CLI usage.
[[autodoc]] lerobot.processor.migrate_policy_normalization.extract_normalization_stats
[[autodoc]] lerobot.processor.migrate_policy_normalization.detect_features_and_norm_modes
[[autodoc]] lerobot.processor.migrate_policy_normalization.remove_normalization_layers
[[autodoc]] lerobot.processor.migrate_policy_normalization.clean_state_dict
[[autodoc]] lerobot.processor.migrate_policy_normalization.load_state_dict_with_missing_key_handling
[[autodoc]] lerobot.processor.migrate_policy_normalization.convert_features_to_policy_features
[[autodoc]] lerobot.processor.migrate_policy_normalization.display_migration_summary_with_warnings
[[autodoc]] lerobot.processor.migrate_policy_normalization.load_model_from_hub
-1
View File
@@ -449,7 +449,6 @@ ignore = [
"src/lerobot/motors/**" = ["D"] "src/lerobot/motors/**" = ["D"]
"src/lerobot/optim/**" = ["D"] "src/lerobot/optim/**" = ["D"]
"src/lerobot/policies/**" = ["D"] "src/lerobot/policies/**" = ["D"]
"src/lerobot/processor/**" = ["D"]
"src/lerobot/rewards/**" = ["D"] "src/lerobot/rewards/**" = ["D"]
"src/lerobot/rl/**" = ["D"] "src/lerobot/rl/**" = ["D"]
"src/lerobot/rollout/**" = ["D"] "src/lerobot/rollout/**" = ["D"]
+13 -26
View File
@@ -14,8 +14,7 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
""" """This script defines processor steps for adding a batch dimension to various components of an environment transition.
This script defines processor steps for adding a batch dimension to various components of an environment transition.
These steps are designed to process actions, observations, and complementary data, making them suitable for batch processing by adding a leading dimension. This is a common requirement before feeding data into a neural network model. These steps are designed to process actions, observations, and complementary data, making them suitable for batch processing by adding a leading dimension. This is a common requirement before feeding data into a neural network model.
""" """
@@ -41,15 +40,13 @@ from .pipeline import (
@dataclass @dataclass
@ProcessorStepRegistry.register(name="to_batch_processor_action") @ProcessorStepRegistry.register(name="to_batch_processor_action")
class AddBatchDimensionActionStep(PolicyActionProcessorStep): class AddBatchDimensionActionStep(PolicyActionProcessorStep):
""" """Processor step to add a batch dimension to a 1D tensor action.
Processor step to add a batch dimension to a 1D tensor action.
This is useful for creating a batch of size 1 from a single action sample. This is useful for creating a batch of size 1 from a single action sample.
""" """
def action(self, action: PolicyAction) -> PolicyAction: def action(self, action: PolicyAction) -> PolicyAction:
""" """Adds a batch dimension to the action if it's a 1D tensor.
Adds a batch dimension to the action if it's a 1D tensor.
Args: Args:
action: The action tensor. action: The action tensor.
@@ -64,8 +61,7 @@ class AddBatchDimensionActionStep(PolicyActionProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Returns the input features unchanged.
Returns the input features unchanged.
Adding a batch dimension does not alter the feature definition. Adding a batch dimension does not alter the feature definition.
@@ -81,8 +77,7 @@ class AddBatchDimensionActionStep(PolicyActionProcessorStep):
@dataclass @dataclass
@ProcessorStepRegistry.register(name="to_batch_processor_observation") @ProcessorStepRegistry.register(name="to_batch_processor_observation")
class AddBatchDimensionObservationStep(ObservationProcessorStep): class AddBatchDimensionObservationStep(ObservationProcessorStep):
""" """Processor step to add a batch dimension to observations.
Processor step to add a batch dimension to observations.
It handles different types of observations: It handles different types of observations:
- State vectors (1D tensors). - State vectors (1D tensors).
@@ -91,8 +86,7 @@ class AddBatchDimensionObservationStep(ObservationProcessorStep):
""" """
def observation(self, observation: dict[str, Tensor]) -> dict[str, Tensor]: def observation(self, observation: dict[str, Tensor]) -> dict[str, Tensor]:
""" """Adds a batch dimension to tensor-based observations in the observation dictionary.
Adds a batch dimension to tensor-based observations in the observation dictionary.
Args: Args:
observation: The observation dictionary. observation: The observation dictionary.
@@ -122,8 +116,7 @@ class AddBatchDimensionObservationStep(ObservationProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Returns the input features unchanged.
Returns the input features unchanged.
Adding a batch dimension does not alter the feature definition. Adding a batch dimension does not alter the feature definition.
@@ -139,8 +132,7 @@ class AddBatchDimensionObservationStep(ObservationProcessorStep):
@dataclass @dataclass
@ProcessorStepRegistry.register(name="to_batch_processor_complementary_data") @ProcessorStepRegistry.register(name="to_batch_processor_complementary_data")
class AddBatchDimensionComplementaryDataStep(ComplementaryDataProcessorStep): class AddBatchDimensionComplementaryDataStep(ComplementaryDataProcessorStep):
""" """Processor step to add a batch dimension to complementary data fields.
Processor step to add a batch dimension to complementary data fields.
Handles specific keys like 'task', 'index', and 'task_index' to make them batched. Handles specific keys like 'task', 'index', and 'task_index' to make them batched.
- 'task' (str) is wrapped in a list. - 'task' (str) is wrapped in a list.
@@ -148,8 +140,7 @@ class AddBatchDimensionComplementaryDataStep(ComplementaryDataProcessorStep):
""" """
def complementary_data(self, complementary_data: dict) -> dict: def complementary_data(self, complementary_data: dict) -> dict:
""" """Adds a batch dimension to specific fields in the complementary data dictionary.
Adds a batch dimension to specific fields in the complementary data dictionary.
Args: Args:
complementary_data: The complementary data dictionary. complementary_data: The complementary data dictionary.
@@ -194,8 +185,7 @@ class AddBatchDimensionComplementaryDataStep(ComplementaryDataProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Returns the input features unchanged.
Returns the input features unchanged.
Adding a batch dimension does not alter the feature definition. Adding a batch dimension does not alter the feature definition.
@@ -211,8 +201,7 @@ class AddBatchDimensionComplementaryDataStep(ComplementaryDataProcessorStep):
@dataclass @dataclass
@ProcessorStepRegistry.register(name="to_batch_processor") @ProcessorStepRegistry.register(name="to_batch_processor")
class AddBatchDimensionProcessorStep(ProcessorStep): class AddBatchDimensionProcessorStep(ProcessorStep):
""" """A composite processor step that adds a batch dimension to the entire environment transition.
A composite processor step that adds a batch dimension to the entire environment transition.
This step combines individual processors for actions, observations, and complementary data This step combines individual processors for actions, observations, and complementary data
to create a batched transition (batch size 1) from a single-instance transition. to create a batched transition (batch size 1) from a single-instance transition.
@@ -236,8 +225,7 @@ class AddBatchDimensionProcessorStep(ProcessorStep):
) )
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
""" """Applies the batching process to all relevant parts of an environment transition.
Applies the batching process to all relevant parts of an environment transition.
Args: Args:
transition: The environment transition to process. transition: The environment transition to process.
@@ -256,8 +244,7 @@ class AddBatchDimensionProcessorStep(ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Returns the input features unchanged.
Returns the input features unchanged.
Adding a batch dimension does not alter the feature definition. Adding a batch dimension does not alter the feature definition.
+62 -47
View File
@@ -34,16 +34,16 @@ def to_tensor(
dtype: torch.dtype | None = torch.float32, dtype: torch.dtype | None = torch.float32,
device: torch.device | str | None = None, device: torch.device | str | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
""" """Convert various data types to PyTorch tensors with configurable options.
Convert various data types to PyTorch tensors with configurable options.
This is a unified tensor conversion function using single dispatch to handle This is a unified tensor conversion function using single dispatch to handle
different input types appropriately. different input types appropriately.
Args: Args:
value: Input value to convert (tensor, array, scalar, sequence, etc.). value (`Any`): Input value to convert (tensor, array, scalar, sequence, etc.).
dtype: Target tensor dtype. If None, preserves original dtype. dtype (`torch.dtype | None`, *optional*, defaults to `torch.float32`): Target tensor dtype. If
device: Target device for the tensor. None, preserves original dtype.
device (`torch.device | str | None`, *optional*): Target device for the tensor.
Returns: Returns:
A PyTorch tensor. A PyTorch tensor.
@@ -137,13 +137,12 @@ def _(value: dict, *, device=None, **kwargs) -> dict:
def from_tensor_to_numpy(x: torch.Tensor | Any) -> np.ndarray | float | int | Any: def from_tensor_to_numpy(x: torch.Tensor | Any) -> np.ndarray | float | int | Any:
""" """Convert a PyTorch tensor to a numpy array or scalar if applicable.
Convert a PyTorch tensor to a numpy array or scalar if applicable.
If the input is not a tensor, it is returned unchanged. If the input is not a tensor, it is returned unchanged.
Args: Args:
x: The input, which can be a tensor or any other type. x (`torch.Tensor | Any`): The input, which can be a tensor or any other type.
Returns: Returns:
A numpy array, a scalar, or the original input. A numpy array, a scalar, or the original input.
@@ -188,17 +187,16 @@ def create_transition(
info: dict[str, Any] | None = None, info: dict[str, Any] | None = None,
complementary_data: dict[str, Any] | None = None, complementary_data: dict[str, Any] | None = None,
) -> EnvTransition: ) -> EnvTransition:
""" """Create an `EnvTransition` dictionary with sensible defaults.
Create an `EnvTransition` dictionary with sensible defaults.
Args: Args:
observation: Observation dictionary. observation (`RobotObservation | None`, *optional*): Observation dictionary.
action: Action dictionary. action (`PolicyAction | RobotAction | None`, *optional*): Action dictionary.
reward: Scalar reward value. reward (`float`, *optional*, defaults to 0.0): Scalar reward value.
done: Episode termination flag. done (`bool`, *optional*, defaults to `False`): Episode termination flag.
truncated: Episode truncation flag. truncated (`bool`, *optional*, defaults to `False`): Episode truncation flag.
info: Additional info dictionary. info (`dict[str, Any] | None`, *optional*): Additional info dictionary.
complementary_data: Complementary data dictionary. complementary_data (`dict[str, Any] | None`, *optional*): Complementary data dictionary.
Returns: Returns:
A complete `EnvTransition` dictionary. A complete `EnvTransition` dictionary.
@@ -217,15 +215,15 @@ def create_transition(
def robot_action_observation_to_transition( def robot_action_observation_to_transition(
action_observation: tuple[RobotAction, RobotObservation], action_observation: tuple[RobotAction, RobotObservation],
) -> EnvTransition: ) -> EnvTransition:
""" """Convert a raw robot action and observation dictionary into a standardized `EnvTransition`.
Convert a raw robot action and observation dictionary into a standardized `EnvTransition`.
Args: Args:
action: The raw action dictionary from a teleoperation device or controller. action_observation (`tuple[RobotAction, RobotObservation]`): A `(action, observation)` tuple, where
observation: The raw observation dictionary from the environment. `action` is the raw action dictionary from a teleoperation device or controller, and
`observation` is the raw observation dictionary from the environment.
Returns: Returns:
An `EnvTransition` containing the formatted observation. An `EnvTransition` containing the formatted action and observation.
""" """
if not isinstance(action_observation, tuple): if not isinstance(action_observation, tuple):
raise ValueError("action_observation should be a tuple type with an action and observation") raise ValueError("action_observation should be a tuple type with an action and observation")
@@ -242,11 +240,10 @@ def robot_action_observation_to_transition(
def robot_action_to_transition(action: RobotAction) -> EnvTransition: def robot_action_to_transition(action: RobotAction) -> EnvTransition:
""" """Convert a raw robot action dictionary into a standardized `EnvTransition`.
Convert a raw robot action dictionary into a standardized `EnvTransition`.
Args: Args:
action: The raw action dictionary from a teleoperation device or controller. action (`RobotAction`): The raw action dictionary from a teleoperation device or controller.
Returns: Returns:
An `EnvTransition` containing the formatted action. An `EnvTransition` containing the formatted action.
@@ -257,11 +254,10 @@ def robot_action_to_transition(action: RobotAction) -> EnvTransition:
def observation_to_transition(observation: RobotObservation) -> EnvTransition: def observation_to_transition(observation: RobotObservation) -> EnvTransition:
""" """Convert a raw robot observation dictionary into a standardized `EnvTransition`.
Convert a raw robot observation dictionary into a standardized `EnvTransition`.
Args: Args:
observation: The raw observation dictionary from the environment. observation (`RobotObservation`): The raw observation dictionary from the environment.
Returns: Returns:
An `EnvTransition` containing the formatted observation. An `EnvTransition` containing the formatted observation.
@@ -272,14 +268,13 @@ def observation_to_transition(observation: RobotObservation) -> EnvTransition:
def transition_to_robot_action(transition: EnvTransition) -> RobotAction: def transition_to_robot_action(transition: EnvTransition) -> RobotAction:
""" """Extract a raw robot action dictionary for a robot from an `EnvTransition`.
Extract a raw robot action dictionary for a robot from an `EnvTransition`.
This function searches for keys in the format "action.*.pos" or "action.*.vel" This function searches for keys in the format "action.*.pos" or "action.*.vel"
and converts them into a flat dictionary suitable for sending to a robot controller. and converts them into a flat dictionary suitable for sending to a robot controller.
Args: Args:
transition: The `EnvTransition` containing the action. transition (`EnvTransition`): The `EnvTransition` containing the action.
Returns: Returns:
A dictionary representing the raw robot action. A dictionary representing the raw robot action.
@@ -294,8 +289,16 @@ def transition_to_robot_action(transition: EnvTransition) -> RobotAction:
def transition_to_policy_action(transition: EnvTransition) -> PolicyAction: def transition_to_policy_action(transition: EnvTransition) -> PolicyAction:
""" """Convert an `EnvTransition` to a `PolicyAction`.
Convert an `EnvTransition` to a `PolicyAction`.
Args:
transition (`EnvTransition`): The `EnvTransition` containing the action.
Returns:
The extracted `PolicyAction`.
Raises:
ValueError: If `transition` is not a dict, or its action is not a `PolicyAction`.
""" """
if not isinstance(transition, dict): if not isinstance(transition, dict):
raise ValueError(f"Transition should be a EnvTransition type (dict) got {type(transition)}") raise ValueError(f"Transition should be a EnvTransition type (dict) got {type(transition)}")
@@ -307,8 +310,16 @@ def transition_to_policy_action(transition: EnvTransition) -> PolicyAction:
def transition_to_observation(transition: EnvTransition) -> RobotObservation: def transition_to_observation(transition: EnvTransition) -> RobotObservation:
""" """Convert an `EnvTransition` to a `RobotObservation`.
Convert an `EnvTransition` to a `RobotObservation`.
Args:
transition (`EnvTransition`): The `EnvTransition` containing the observation.
Returns:
The extracted `RobotObservation`.
Raises:
ValueError: If `transition` is not a dict, or its observation is not a dict.
""" """
if not isinstance(transition, dict): if not isinstance(transition, dict):
raise ValueError(f"Transition should be a EnvTransition type (dict) got {type(transition)}") raise ValueError(f"Transition should be a EnvTransition type (dict) got {type(transition)}")
@@ -320,8 +331,16 @@ def transition_to_observation(transition: EnvTransition) -> RobotObservation:
def policy_action_to_transition(action: PolicyAction) -> EnvTransition: def policy_action_to_transition(action: PolicyAction) -> EnvTransition:
""" """Convert a `PolicyAction` to an `EnvTransition`.
Convert a `PolicyAction` to an `EnvTransition`.
Args:
action (`PolicyAction`): The `PolicyAction` to wrap.
Returns:
An `EnvTransition` containing the formatted action.
Raises:
ValueError: If `action` is not a `PolicyAction`.
""" """
if not isinstance(action, PolicyAction): if not isinstance(action, PolicyAction):
raise ValueError(f"Action should be a PolicyAction type got {type(action)}") raise ValueError(f"Action should be a PolicyAction type got {type(action)}")
@@ -329,14 +348,13 @@ def policy_action_to_transition(action: PolicyAction) -> EnvTransition:
def batch_to_transition(batch: dict[str, Any]) -> EnvTransition: def batch_to_transition(batch: dict[str, Any]) -> EnvTransition:
""" """Convert a batch dictionary from a dataset/dataloader into an `EnvTransition`.
Convert a batch dictionary from a dataset/dataloader into an `EnvTransition`.
This function maps recognized keys from a batch to the `EnvTransition` structure, This function maps recognized keys from a batch to the `EnvTransition` structure,
filling in missing keys with sensible defaults. filling in missing keys with sensible defaults.
Args: Args:
batch: A batch dictionary. batch (`dict[str, Any]`): A batch dictionary.
Returns: Returns:
An `EnvTransition` dictionary. An `EnvTransition` dictionary.
@@ -344,7 +362,6 @@ def batch_to_transition(batch: dict[str, Any]) -> EnvTransition:
Raises: Raises:
ValueError: If the input is not a dictionary. ValueError: If the input is not a dictionary.
""" """
# Validate input type. # Validate input type.
if not isinstance(batch, dict): if not isinstance(batch, dict):
raise ValueError(f"EnvTransition must be a dictionary. Got {type(batch).__name__}") raise ValueError(f"EnvTransition must be a dictionary. Got {type(batch).__name__}")
@@ -369,13 +386,12 @@ def batch_to_transition(batch: dict[str, Any]) -> EnvTransition:
def transition_to_batch(transition: EnvTransition) -> dict[str, Any]: def transition_to_batch(transition: EnvTransition) -> dict[str, Any]:
""" """Convert an `EnvTransition` back to the canonical batch format used in LeRobot.
Convert an `EnvTransition` back to the canonical batch format used in LeRobot.
This is the inverse of `batch_to_transition`. This is the inverse of `batch_to_transition`.
Args: Args:
transition: The `EnvTransition` to convert. transition (`EnvTransition`): The `EnvTransition` to convert.
Returns: Returns:
A batch dictionary with canonical LeRobot field names. A batch dictionary with canonical LeRobot field names.
@@ -405,13 +421,12 @@ def transition_to_batch(transition: EnvTransition) -> dict[str, Any]:
def identity_transition(transition: EnvTransition) -> EnvTransition: def identity_transition(transition: EnvTransition) -> EnvTransition:
""" """An identity function for transitions, returning the input unchanged.
An identity function for transitions, returning the input unchanged.
Useful as a default or placeholder in processing pipelines. Useful as a default or placeholder in processing pipelines.
Args: Args:
tr: An `EnvTransition`. transition (`EnvTransition`): An `EnvTransition`.
Returns: Returns:
The same `EnvTransition`. The same `EnvTransition`.
@@ -25,8 +25,7 @@ from .pipeline import ActionProcessorStep, ProcessorStepRegistry, RobotActionPro
@ProcessorStepRegistry.register("map_tensor_to_delta_action_dict") @ProcessorStepRegistry.register("map_tensor_to_delta_action_dict")
@dataclass @dataclass
class MapTensorToDeltaActionDictStep(ActionProcessorStep): class MapTensorToDeltaActionDictStep(ActionProcessorStep):
""" """Maps a flat action tensor from a policy to a structured delta action dictionary.
Maps a flat action tensor from a policy to a structured delta action dictionary.
This step is typically used after a policy outputs a continuous action vector. This step is typically used after a policy outputs a continuous action vector.
It decomposes the vector into named components for delta movements of the It decomposes the vector into named components for delta movements of the
@@ -39,6 +38,18 @@ class MapTensorToDeltaActionDictStep(ActionProcessorStep):
use_gripper: bool = True use_gripper: bool = True
def action(self, action: PolicyAction) -> RobotAction: def action(self, action: PolicyAction) -> RobotAction:
"""Split a flat policy action tensor into a named delta-movement dict.
Args:
action: A `PolicyAction` tensor of at least 3 elements (x, y, z), plus a 4th gripper
element if `use_gripper` is `True`.
Returns:
A dict with `delta_x`/`delta_y`/`delta_z` (and `gripper`, if enabled).
Raises:
ValueError: If `action` is not a `PolicyAction`.
"""
if not isinstance(action, PolicyAction): if not isinstance(action, PolicyAction):
raise ValueError("Only PolicyAction is supported for this processor") raise ValueError("Only PolicyAction is supported for this processor")
@@ -58,6 +69,7 @@ class MapTensorToDeltaActionDictStep(ActionProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. Adds the `delta_x`/`delta_y`/`delta_z` (and `gripper`) action features."""
for axis in ["x", "y", "z"]: for axis in ["x", "y", "z"]:
features[PipelineFeatureType.ACTION][f"delta_{axis}"] = PolicyFeature( features[PipelineFeatureType.ACTION][f"delta_{axis}"] = PolicyFeature(
type=FeatureType.ACTION, shape=(1,) type=FeatureType.ACTION, shape=(1,)
@@ -73,8 +85,7 @@ class MapTensorToDeltaActionDictStep(ActionProcessorStep):
@ProcessorStepRegistry.register("map_delta_action_to_robot_action") @ProcessorStepRegistry.register("map_delta_action_to_robot_action")
@dataclass @dataclass
class MapDeltaActionToRobotActionStep(RobotActionProcessorStep): class MapDeltaActionToRobotActionStep(RobotActionProcessorStep):
""" """Maps delta actions from teleoperators to robot target actions for inverse kinematics.
Maps delta actions from teleoperators to robot target actions for inverse kinematics.
This step converts a dictionary of delta movements (e.g., from a gamepad) This step converts a dictionary of delta movements (e.g., from a gamepad)
into a target action format that includes an "enabled" flag and target into a target action format that includes an "enabled" flag and target
@@ -91,6 +102,15 @@ class MapDeltaActionToRobotActionStep(RobotActionProcessorStep):
noise_threshold: float = 1e-3 # 1 mm threshold to filter out noise noise_threshold: float = 1e-3 # 1 mm threshold to filter out noise
def action(self, action: RobotAction) -> RobotAction: def action(self, action: RobotAction) -> RobotAction:
"""Convert a delta-movement dict into a robot target-action dict for inverse kinematics.
Args:
action: A dict with `delta_x`/`delta_y`/`delta_z` and `gripper` keys.
Returns:
A dict with `enabled`, scaled `target_x`/`target_y`/`target_z`, zeroed `target_wx`/`target_wy`/
`target_wz` (rotation isn't supported by delta teleoperators), and `gripper_vel`.
"""
# NOTE (maractingi): Action can be a dict from the teleop_devices or a tensor from the policy # NOTE (maractingi): Action can be a dict from the teleop_devices or a tensor from the policy
# TODO (maractingi): changing this target_xyz naming convention from the teleop_devices # TODO (maractingi): changing this target_xyz naming convention from the teleop_devices
delta_x = action.pop("delta_x") delta_x = action.pop("delta_x")
@@ -131,6 +151,7 @@ class MapDeltaActionToRobotActionStep(RobotActionProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. Replaces the delta-action features with robot target-action features."""
for axis in ["x", "y", "z"]: for axis in ["x", "y", "z"]:
features[PipelineFeatureType.ACTION].pop(f"delta_{axis}", None) features[PipelineFeatureType.ACTION].pop(f"delta_{axis}", None)
features[PipelineFeatureType.ACTION].pop("gripper", None) features[PipelineFeatureType.ACTION].pop("gripper", None)
+11 -17
View File
@@ -14,9 +14,9 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
""" """This script defines a processor step for moving environment transition data to a specific torch device.
This script defines a processor step for moving environment transition data to a specific torch device and casting
its floating-point precision. It also optionally casts data to a specified floating-point precision.
""" """
from dataclasses import dataclass from dataclasses import dataclass
@@ -34,11 +34,10 @@ from .pipeline import ProcessorStep, ProcessorStepRegistry
@ProcessorStepRegistry.register("device_processor") @ProcessorStepRegistry.register("device_processor")
@dataclass @dataclass
class DeviceProcessorStep(ProcessorStep): class DeviceProcessorStep(ProcessorStep):
""" """Processor step to move all tensors within an `EnvTransition` to a specified device.
Processor step to move all tensors within an `EnvTransition` to a specified device and optionally cast their
floating-point data type.
This is crucial for preparing data for model training or inference on hardware like GPUs. Optionally casts their floating-point data type too. This is crucial for preparing data for model
training or inference on hardware like GPUs.
**Attributes**: **Attributes**:
- **device** (`str`) -- The target device for tensors (e.g., "cpu", "cuda", "cuda:0"). - **device** (`str`) -- The target device for tensors (e.g., "cpu", "cuda", "cuda:0").
@@ -60,8 +59,7 @@ class DeviceProcessorStep(ProcessorStep):
} }
def __post_init__(self): def __post_init__(self):
""" """Initializes the processor by converting string configurations to torch objects.
Initializes the processor by converting string configurations to torch objects.
This method sets up the `torch.device`, determines if transfers can be non-blocking, and validates the This method sets up the `torch.device`, determines if transfers can be non-blocking, and validates the
`float_dtype` string, converting it to a `torch.dtype` object. `float_dtype` string, converting it to a `torch.dtype` object.
@@ -82,8 +80,7 @@ class DeviceProcessorStep(ProcessorStep):
self._target_float_dtype = None self._target_float_dtype = None
def _process_tensor(self, tensor: torch.Tensor) -> torch.Tensor: def _process_tensor(self, tensor: torch.Tensor) -> torch.Tensor:
""" """Moves a single tensor to the target device and casts its dtype.
Moves a single tensor to the target device and casts its dtype.
Handles multi-GPU scenarios by not moving a tensor if it's already on a different CUDA device than Handles multi-GPU scenarios by not moving a tensor if it's already on a different CUDA device than
the target, which is useful when using frameworks like Accelerate. the target, which is useful when using frameworks like Accelerate.
@@ -120,8 +117,7 @@ class DeviceProcessorStep(ProcessorStep):
return tensor return tensor
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
""" """Applies device and dtype conversion to all tensors in an environment transition.
Applies device and dtype conversion to all tensors in an environment transition.
It iterates through the transition, finds all `torch.Tensor` objects (including those nested in It iterates through the transition, finds all `torch.Tensor` objects (including those nested in
dictionaries like `observation`), and processes them. dictionaries like `observation`), and processes them.
@@ -169,8 +165,7 @@ class DeviceProcessorStep(ProcessorStep):
return new_transition return new_transition
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
""" """Returns the serializable configuration of the processor.
Returns the serializable configuration of the processor.
Returns: Returns:
A dictionary containing the device and float_dtype settings. A dictionary containing the device and float_dtype settings.
@@ -180,8 +175,7 @@ class DeviceProcessorStep(ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Returns the input features unchanged.
Returns the input features unchanged.
Device and dtype transformations do not alter the fundamental definition of the features (e.g., shape). Device and dtype transformations do not alter the fundamental definition of the features (e.g., shape).
+9 -16
View File
@@ -26,8 +26,7 @@ from .pipeline import ObservationProcessorStep, ProcessorStepRegistry
@dataclass @dataclass
@ProcessorStepRegistry.register(name="libero_processor") @ProcessorStepRegistry.register(name="libero_processor")
class LiberoProcessorStep(ObservationProcessorStep): class LiberoProcessorStep(ObservationProcessorStep):
""" """Processes LIBERO observations into the LeRobot format.
Processes LIBERO observations into the LeRobot format.
This step handles the specific observation structure from LIBERO environments, This step handles the specific observation structure from LIBERO environments,
which includes nested robot_state dictionaries and image observations. which includes nested robot_state dictionaries and image observations.
@@ -47,9 +46,7 @@ class LiberoProcessorStep(ObservationProcessorStep):
""" """
def _process_observation(self, observation): def _process_observation(self, observation):
""" """Processes both image and robot_state observations from LIBERO."""
Processes both image and robot_state observations from LIBERO.
"""
processed_obs = observation.copy() processed_obs = observation.copy()
for key in list(processed_obs.keys()): for key in list(processed_obs.keys()):
if key.startswith(f"{OBS_IMAGES}."): if key.startswith(f"{OBS_IMAGES}."):
@@ -85,9 +82,7 @@ class LiberoProcessorStep(ObservationProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Transforms feature keys from the LIBERO format to the LeRobot standard."""
Transforms feature keys from the LIBERO format to the LeRobot standard.
"""
new_features: dict[PipelineFeatureType, dict[str, PolicyFeature]] = {} new_features: dict[PipelineFeatureType, dict[str, PolicyFeature]] = {}
# copy over non-STATE features # copy over non-STATE features
@@ -109,11 +104,12 @@ class LiberoProcessorStep(ObservationProcessorStep):
return new_features return new_features
def observation(self, observation): def observation(self, observation):
"""See [`~processor.ObservationProcessorStep.observation`]. Delegates to `_process_observation`."""
return self._process_observation(observation) return self._process_observation(observation)
def _quat2axisangle(self, quat: torch.Tensor) -> torch.Tensor: def _quat2axisangle(self, quat: torch.Tensor) -> torch.Tensor:
""" """Convert batched quaternions to axis-angle format.
Convert batched quaternions to axis-angle format.
Only accepts torch tensors of shape (B, 4). Only accepts torch tensors of shape (B, 4).
Args: Args:
@@ -126,7 +122,6 @@ class LiberoProcessorStep(ObservationProcessorStep):
TypeError: if input is not a torch tensor TypeError: if input is not a torch tensor
ValueError: if shape is not (B, 4) ValueError: if shape is not (B, 4)
""" """
if not isinstance(quat, torch.Tensor): if not isinstance(quat, torch.Tensor):
raise TypeError(f"_quat2axisangle expected a torch.Tensor, got {type(quat)}") raise TypeError(f"_quat2axisangle expected a torch.Tensor, got {type(quat)}")
@@ -156,8 +151,7 @@ class LiberoProcessorStep(ObservationProcessorStep):
@dataclass @dataclass
@ProcessorStepRegistry.register(name="isaaclab_arena_processor") @ProcessorStepRegistry.register(name="isaaclab_arena_processor")
class IsaaclabArenaProcessorStep(ObservationProcessorStep): class IsaaclabArenaProcessorStep(ObservationProcessorStep):
""" """Processes IsaacLab Arena observations into LeRobot format.
Processes IsaacLab Arena observations into LeRobot format.
**State Processing:** **State Processing:**
- Extracts state components from obs["policy"] based on `state_keys`. - Extracts state components from obs["policy"] based on `state_keys`.
@@ -176,9 +170,7 @@ class IsaaclabArenaProcessorStep(ObservationProcessorStep):
camera_keys: tuple[str, ...] camera_keys: tuple[str, ...]
def _process_observation(self, observation): def _process_observation(self, observation):
""" """Processes both image and policy state observations from IsaacLab Arena."""
Processes both image and policy state observations from IsaacLab Arena.
"""
processed_obs = {} processed_obs = {}
if f"{OBS_STR}.camera_obs" in observation: if f"{OBS_STR}.camera_obs" in observation:
@@ -225,4 +217,5 @@ class IsaaclabArenaProcessorStep(ObservationProcessorStep):
return features return features
def observation(self, observation): def observation(self, observation):
"""See [`~processor.ObservationProcessorStep.observation`]. Delegates to `_process_observation`."""
return self._process_observation(observation) return self._process_observation(observation)
+29 -6
View File
@@ -46,6 +46,11 @@ from .rename_processor import RenameObservationsProcessorStep
def make_default_teleop_action_processor() -> RobotProcessorPipeline[ def make_default_teleop_action_processor() -> RobotProcessorPipeline[
tuple[RobotAction, RobotObservation], RobotAction tuple[RobotAction, RobotObservation], RobotAction
]: ]:
"""Build a no-op teleoperator-action pipeline (an `IdentityProcessorStep`).
Returns:
`RobotProcessorPipeline[tuple[RobotAction, RobotObservation], RobotAction]`: The pipeline.
"""
teleop_action_processor = RobotProcessorPipeline[tuple[RobotAction, RobotObservation], RobotAction]( teleop_action_processor = RobotProcessorPipeline[tuple[RobotAction, RobotObservation], RobotAction](
steps=[IdentityProcessorStep()], steps=[IdentityProcessorStep()],
to_transition=robot_action_observation_to_transition, to_transition=robot_action_observation_to_transition,
@@ -57,6 +62,11 @@ def make_default_teleop_action_processor() -> RobotProcessorPipeline[
def make_default_robot_action_processor() -> RobotProcessorPipeline[ def make_default_robot_action_processor() -> RobotProcessorPipeline[
tuple[RobotAction, RobotObservation], RobotAction tuple[RobotAction, RobotObservation], RobotAction
]: ]:
"""Build a no-op robot-action pipeline (an `IdentityProcessorStep`).
Returns:
`RobotProcessorPipeline[tuple[RobotAction, RobotObservation], RobotAction]`: The pipeline.
"""
robot_action_processor = RobotProcessorPipeline[tuple[RobotAction, RobotObservation], RobotAction]( robot_action_processor = RobotProcessorPipeline[tuple[RobotAction, RobotObservation], RobotAction](
steps=[IdentityProcessorStep()], steps=[IdentityProcessorStep()],
to_transition=robot_action_observation_to_transition, to_transition=robot_action_observation_to_transition,
@@ -66,6 +76,11 @@ def make_default_robot_action_processor() -> RobotProcessorPipeline[
def make_default_robot_observation_processor() -> RobotProcessorPipeline[RobotObservation, RobotObservation]: def make_default_robot_observation_processor() -> RobotProcessorPipeline[RobotObservation, RobotObservation]:
"""Build a no-op robot-observation pipeline (an `IdentityProcessorStep`).
Returns:
`RobotProcessorPipeline[RobotObservation, RobotObservation]`: The pipeline.
"""
robot_observation_processor = RobotProcessorPipeline[RobotObservation, RobotObservation]( robot_observation_processor = RobotProcessorPipeline[RobotObservation, RobotObservation](
steps=[IdentityProcessorStep()], steps=[IdentityProcessorStep()],
to_transition=observation_to_transition, to_transition=observation_to_transition,
@@ -75,6 +90,11 @@ def make_default_robot_observation_processor() -> RobotProcessorPipeline[RobotOb
def make_default_processors(): def make_default_processors():
"""Build the three no-op default processors: teleop-action, robot-action, and robot-observation.
Returns:
A `(teleop_action_processor, robot_action_processor, robot_observation_processor)` tuple.
"""
teleop_action_processor = make_default_teleop_action_processor() teleop_action_processor = make_default_teleop_action_processor()
robot_action_processor = make_default_robot_action_processor() robot_action_processor = make_default_robot_action_processor()
robot_observation_processor = make_default_robot_observation_processor() robot_observation_processor = make_default_robot_observation_processor()
@@ -106,11 +126,13 @@ def make_default_policy_processor_steps(
"""Construct the canonical policy processor steps from a policy config. """Construct the canonical policy processor steps from a policy config.
Args: Args:
config: A `PreTrainedConfig` providing `device`, `input_features`, config (`PreTrainedConfig`): A `PreTrainedConfig` providing `device`, `input_features`,
`output_features` and `normalization_mapping`. `output_features` and `normalization_mapping`.
dataset_stats: Dataset statistics used for (un)normalization. dataset_stats (`dict[str, dict[str, torch.Tensor]] | None`, *optional*): Dataset statistics used
normalizer_device: Device passed to `NormalizerProcessorStep` (some policies pin for (un)normalization.
their normalization stats to the policy device; most leave it unset). normalizer_device (`torch.device | str | None`, *optional*): Device passed to
`NormalizerProcessorStep` (some policies pin their normalization stats to the policy device;
most leave it unset).
""" """
return DefaultPolicyProcessorSteps( return DefaultPolicyProcessorSteps(
rename_observations=RenameObservationsProcessorStep(rename_map={}), rename_observations=RenameObservationsProcessorStep(rename_map={}),
@@ -164,8 +186,9 @@ def make_default_pre_post_processors(
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]], PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
PolicyProcessorPipeline[PolicyAction, PolicyAction], PolicyProcessorPipeline[PolicyAction, PolicyAction],
]: ]:
"""The pure-scaffold policy pipeline pair: Rename -> Batch -> Device -> Normalize, """The pure-scaffold policy pipeline pair: Rename -> Batch -> Device -> Normalize.
and Unnormalize -> Device(cpu). Policies with custom steps or a different step order
And Unnormalize -> Device(cpu). Policies with custom steps or a different step order
compose `make_default_policy_processor_steps` themselves instead. compose `make_default_policy_processor_steps` themselves instead.
""" """
s = make_default_policy_processor_steps(config, dataset_stats, normalizer_device=normalizer_device) s = make_default_policy_processor_steps(config, dataset_stats, normalizer_device=normalizer_device)
@@ -27,8 +27,7 @@ from .pipeline import ActionProcessorStep, ProcessorStep, ProcessorStepRegistry
@ProcessorStepRegistry.register("torch2numpy_action_processor") @ProcessorStepRegistry.register("torch2numpy_action_processor")
@dataclass @dataclass
class Torch2NumpyActionProcessorStep(ActionProcessorStep): class Torch2NumpyActionProcessorStep(ActionProcessorStep):
""" """Converts a PyTorch tensor action to a NumPy array.
Converts a PyTorch tensor action to a NumPy array.
This step is useful when the output of a policy (typically a torch.Tensor) This step is useful when the output of a policy (typically a torch.Tensor)
needs to be passed to an environment or component that expects a NumPy array. needs to be passed to an environment or component that expects a NumPy array.
@@ -41,6 +40,11 @@ class Torch2NumpyActionProcessorStep(ActionProcessorStep):
squeeze_batch_dim: bool = True squeeze_batch_dim: bool = True
def action(self, action: PolicyAction) -> EnvAction: def action(self, action: PolicyAction) -> EnvAction:
"""Convert `action` to a NumPy array, squeezing a size-1 batch dimension if `squeeze_batch_dim`.
Raises:
TypeError: If `action` is not a `PolicyAction`.
"""
if not isinstance(action, PolicyAction): if not isinstance(action, PolicyAction):
raise TypeError( raise TypeError(
f"Expected PolicyAction or None, got {type(action).__name__}. " f"Expected PolicyAction or None, got {type(action).__name__}. "
@@ -64,6 +68,7 @@ class Torch2NumpyActionProcessorStep(ActionProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. A dtype conversion; features are unchanged."""
return features return features
@@ -99,4 +104,5 @@ class Numpy2TorchActionProcessorStep(ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. A dtype conversion; features are unchanged."""
return features return features
+32 -48
View File
@@ -47,8 +47,7 @@ TELEOP_ACTION_KEY = "teleop_action"
@runtime_checkable @runtime_checkable
class HasTeleopEvents(Protocol): class HasTeleopEvents(Protocol):
""" """Minimal protocol for objects that provide teleoperation events.
Minimal protocol for objects that provide teleoperation events.
This protocol defines the `get_teleop_events()` method, allowing processor This protocol defines the `get_teleop_events()` method, allowing processor
steps to interact with teleoperators that support event-based controls steps to interact with teleoperators that support event-based controls
@@ -57,8 +56,7 @@ class HasTeleopEvents(Protocol):
""" """
def get_teleop_events(self) -> dict[str, Any]: def get_teleop_events(self) -> dict[str, Any]:
""" """Get extra control events from the teleoperator.
Get extra control events from the teleoperator.
Returns: Returns:
A dictionary containing control events such as: A dictionary containing control events such as:
@@ -75,8 +73,7 @@ TeleopWithEvents = TypeVar("TeleopWithEvents", bound="Teleoperator")
def _check_teleop_with_events(teleop: "Teleoperator") -> None: def _check_teleop_with_events(teleop: "Teleoperator") -> None:
""" """Runtime check that a teleoperator implements the `HasTeleopEvents` protocol.
Runtime check that a teleoperator implements the `HasTeleopEvents` protocol.
Args: Args:
teleop: The teleoperator instance to check. teleop: The teleoperator instance to check.
@@ -94,8 +91,7 @@ def _check_teleop_with_events(teleop: "Teleoperator") -> None:
@ProcessorStepRegistry.register("add_teleop_action_as_complementary_data") @ProcessorStepRegistry.register("add_teleop_action_as_complementary_data")
@dataclass @dataclass
class AddTeleopActionAsComplimentaryDataStep(ComplementaryDataProcessorStep): class AddTeleopActionAsComplimentaryDataStep(ComplementaryDataProcessorStep):
""" """Adds the raw action from a teleoperator to the transition's complementary data.
Adds the raw action from a teleoperator to the transition's complementary data.
This is useful for human-in-the-loop scenarios where the human's input needs to This is useful for human-in-the-loop scenarios where the human's input needs to
be available to downstream processors, for example, to override a policy's action be available to downstream processors, for example, to override a policy's action
@@ -108,8 +104,7 @@ class AddTeleopActionAsComplimentaryDataStep(ComplementaryDataProcessorStep):
teleop_device: "Teleoperator" teleop_device: "Teleoperator"
def complementary_data(self, complementary_data: dict) -> dict: def complementary_data(self, complementary_data: dict) -> dict:
""" """Retrieves the teleoperator's action and adds it to the complementary data.
Retrieves the teleoperator's action and adds it to the complementary data.
Args: Args:
complementary_data: The incoming complementary data dictionary. complementary_data: The incoming complementary data dictionary.
@@ -125,14 +120,14 @@ class AddTeleopActionAsComplimentaryDataStep(ComplementaryDataProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. Complementary data isn't tracked in features; unchanged."""
return features return features
@ProcessorStepRegistry.register("add_teleop_action_as_info") @ProcessorStepRegistry.register("add_teleop_action_as_info")
@dataclass @dataclass
class AddTeleopEventsAsInfoStep(InfoProcessorStep): class AddTeleopEventsAsInfoStep(InfoProcessorStep):
""" """Adds teleoperator control events (e.g., terminate, success) to the transition's info.
Adds teleoperator control events (e.g., terminate, success) to the transition's info.
This step extracts control events from teleoperators that support event-based This step extracts control events from teleoperators that support event-based
interaction, making these signals available to other parts of the system. interaction, making these signals available to other parts of the system.
@@ -149,8 +144,7 @@ class AddTeleopEventsAsInfoStep(InfoProcessorStep):
_check_teleop_with_events(self.teleop_device) _check_teleop_with_events(self.teleop_device)
def info(self, info: dict) -> dict: def info(self, info: dict) -> dict:
""" """Retrieves teleoperator events and updates the info dictionary.
Retrieves teleoperator events and updates the info dictionary.
Args: Args:
info: The incoming info dictionary. info: The incoming info dictionary.
@@ -167,14 +161,14 @@ class AddTeleopEventsAsInfoStep(InfoProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. Info isn't tracked in features; unchanged."""
return features return features
@ProcessorStepRegistry.register("image_crop_resize_processor") @ProcessorStepRegistry.register("image_crop_resize_processor")
@dataclass @dataclass
class ImageCropResizeProcessorStep(ObservationProcessorStep): class ImageCropResizeProcessorStep(ObservationProcessorStep):
""" """Crops and/or resizes image observations.
Crops and/or resizes image observations.
This step iterates through all image keys in an observation dictionary and applies This step iterates through all image keys in an observation dictionary and applies
the specified transformations. It handles device placement, moving tensors to the the specified transformations. It handles device placement, moving tensors to the
@@ -190,8 +184,7 @@ class ImageCropResizeProcessorStep(ObservationProcessorStep):
resize_size: tuple[int, int] | None = None resize_size: tuple[int, int] | None = None
def observation(self, observation: dict) -> dict: def observation(self, observation: dict) -> dict:
""" """Applies cropping and resizing to all images in the observation dictionary.
Applies cropping and resizing to all images in the observation dictionary.
Args: Args:
observation: The observation dictionary, potentially containing image tensors. observation: The observation dictionary, potentially containing image tensors.
@@ -226,8 +219,7 @@ class ImageCropResizeProcessorStep(ObservationProcessorStep):
return new_observation return new_observation
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
""" """Returns the configuration of the step for serialization.
Returns the configuration of the step for serialization.
Returns: Returns:
A dictionary with the crop parameters and resize dimensions. A dictionary with the crop parameters and resize dimensions.
@@ -240,8 +232,7 @@ class ImageCropResizeProcessorStep(ObservationProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Updates the image feature shapes in the policy features dictionary if resizing is applied.
Updates the image feature shapes in the policy features dictionary if resizing is applied.
Args: Args:
features: The policy features dictionary. features: The policy features dictionary.
@@ -264,8 +255,7 @@ class ImageCropResizeProcessorStep(ObservationProcessorStep):
@dataclass @dataclass
@ProcessorStepRegistry.register("time_limit_processor") @ProcessorStepRegistry.register("time_limit_processor")
class TimeLimitProcessorStep(TruncatedProcessorStep): class TimeLimitProcessorStep(TruncatedProcessorStep):
""" """Tracks episode steps and enforces a time limit by truncating the episode.
Tracks episode steps and enforces a time limit by truncating the episode.
**Attributes**: **Attributes**:
- **max_episode_steps** (`int`) -- The maximum number of steps allowed per episode. - **max_episode_steps** (`int`) -- The maximum number of steps allowed per episode.
@@ -276,8 +266,7 @@ class TimeLimitProcessorStep(TruncatedProcessorStep):
current_step: int = 0 current_step: int = 0
def truncated(self, truncated: bool) -> bool: def truncated(self, truncated: bool) -> bool:
""" """Increments the step counter and sets the truncated flag if the time limit is reached.
Increments the step counter and sets the truncated flag if the time limit is reached.
Args: Args:
truncated: The incoming truncated flag. truncated: The incoming truncated flag.
@@ -292,8 +281,7 @@ class TimeLimitProcessorStep(TruncatedProcessorStep):
return truncated return truncated
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
""" """Returns the configuration of the step for serialization.
Returns the configuration of the step for serialization.
Returns: Returns:
A dictionary containing the `max_episode_steps`. A dictionary containing the `max_episode_steps`.
@@ -309,13 +297,13 @@ class TimeLimitProcessorStep(TruncatedProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. The truncated flag isn't tracked in features; unchanged."""
return features return features
@ProcessorStepRegistry.register("gym_hil_adapter_processor") @ProcessorStepRegistry.register("gym_hil_adapter_processor")
class GymHILAdapterProcessorStep(ProcessorStep): class GymHILAdapterProcessorStep(ProcessorStep):
""" """Adapts the output of the `gym-hil` environment to the format expected by `lerobot` processors.
Adapts the output of the `gym-hil` environment to the format expected by `lerobot` processors.
This step normalizes the `transition` object by: This step normalizes the `transition` object by:
1. Copying `teleop_action` from `info` to `complementary_data`. 1. Copying `teleop_action` from `info` to `complementary_data`.
@@ -324,6 +312,7 @@ class GymHILAdapterProcessorStep(ProcessorStep):
""" """
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
"""See [`~processor.ProcessorStep.__call__`]. Performs the key copies described in the class docstring."""
info = transition.get(TransitionKey.INFO, {}) info = transition.get(TransitionKey.INFO, {})
complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
@@ -344,14 +333,14 @@ class GymHILAdapterProcessorStep(ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. A key-copying step; features are unchanged."""
return features return features
@dataclass @dataclass
@ProcessorStepRegistry.register("gripper_penalty_processor") @ProcessorStepRegistry.register("gripper_penalty_processor")
class GripperPenaltyProcessorStep(ProcessorStep): class GripperPenaltyProcessorStep(ProcessorStep):
""" """Applies a small per-transition cost on the discrete gripper action.
Applies a small per-transition cost on the discrete gripper action.
Fires only when the commanded action would actually transition the gripper Fires only when the commanded action would actually transition the gripper
from one extreme to the other (close-while-open or open-while-closed). from one extreme to the other (close-while-open or open-while-closed).
@@ -371,8 +360,7 @@ class GripperPenaltyProcessorStep(ProcessorStep):
closed_threshold: float = 0.9 closed_threshold: float = 0.9
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
""" """Calculates the gripper penalty and adds it to the complementary data.
Calculates the gripper penalty and adds it to the complementary data.
Args: Args:
transition: The incoming environment transition. transition: The incoming environment transition.
@@ -422,8 +410,7 @@ class GripperPenaltyProcessorStep(ProcessorStep):
return new_transition return new_transition
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
""" """Returns the configuration of the step for serialization.
Returns the configuration of the step for serialization.
Returns: Returns:
A dictionary containing the penalty value, max gripper position, A dictionary containing the penalty value, max gripper position,
@@ -443,14 +430,14 @@ class GripperPenaltyProcessorStep(ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. The penalty lives in complementary data, not features; unchanged."""
return features return features
@dataclass @dataclass
@ProcessorStepRegistry.register("intervention_action_processor") @ProcessorStepRegistry.register("intervention_action_processor")
class InterventionActionProcessorStep(ProcessorStep): class InterventionActionProcessorStep(ProcessorStep):
""" """Handles human intervention, overriding policy actions and managing episode termination.
Handles human intervention, overriding policy actions and managing episode termination.
When an intervention is detected (via teleoperator events in the `info` dict), When an intervention is detected (via teleoperator events in the `info` dict),
this step replaces the policy's action with the human's teleoperated action. this step replaces the policy's action with the human's teleoperated action.
@@ -466,8 +453,7 @@ class InterventionActionProcessorStep(ProcessorStep):
terminate_on_success: bool = True terminate_on_success: bool = True
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
""" """Processes the transition to handle interventions.
Processes the transition to handle interventions.
Args: Args:
transition: The incoming environment transition. transition: The incoming environment transition.
@@ -531,8 +517,7 @@ class InterventionActionProcessorStep(ProcessorStep):
return new_transition return new_transition
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
""" """Returns the configuration of the step for serialization.
Returns the configuration of the step for serialization.
Returns: Returns:
A dictionary containing the step's configuration attributes. A dictionary containing the step's configuration attributes.
@@ -545,14 +530,14 @@ class InterventionActionProcessorStep(ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. Overrides the action value, not its shape/type; unchanged."""
return features return features
@dataclass @dataclass
@ProcessorStepRegistry.register("reward_classifier_processor") @ProcessorStepRegistry.register("reward_classifier_processor")
class RewardClassifierProcessorStep(ProcessorStep): class RewardClassifierProcessorStep(ProcessorStep):
""" """Applies a pretrained reward classifier to image observations to predict success.
Applies a pretrained reward classifier to image observations to predict success.
This step uses a model to determine if the current state is successful, updating This step uses a model to determine if the current state is successful, updating
the reward and potentially terminating the episode. the reward and potentially terminating the episode.
@@ -584,8 +569,7 @@ class RewardClassifierProcessorStep(ProcessorStep):
self.reward_classifier.eval() self.reward_classifier.eval()
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
""" """Processes a transition, applying the reward classifier to its image observations.
Processes a transition, applying the reward classifier to its image observations.
Args: Args:
transition: The incoming environment transition. transition: The incoming environment transition.
@@ -633,8 +617,7 @@ class RewardClassifierProcessorStep(ProcessorStep):
return new_transition return new_transition
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
""" """Returns the configuration of the step for serialization.
Returns the configuration of the step for serialization.
Returns: Returns:
A dictionary containing the step's configuration attributes. A dictionary containing the step's configuration attributes.
@@ -649,4 +632,5 @@ class RewardClassifierProcessorStep(ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. Updates reward/done, not features; unchanged."""
return features return features
@@ -14,11 +14,9 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
""" """A generic script to migrate LeRobot policies with built-in normalization layers.
A generic script to migrate LeRobot policies with built-in normalization layers to the new
pipeline-based processor system.
This script performs the following steps: Migrates them to the new pipeline-based processor system. This script performs the following steps:
1. Loads a pretrained policy model and its configuration from a local path or the 1. Loads a pretrained policy model and its configuration from a local path or the
Hugging Face Hub. Hugging Face Hub.
2. Scans the model's state dictionary to extract normalization statistics (e.g., mean, 2. Scans the model's state dictionary to extract normalization statistics (e.g., mean,
@@ -63,15 +61,14 @@ from lerobot.utils.constants import ACTION
def extract_normalization_stats(state_dict: dict[str, torch.Tensor]) -> dict[str, dict[str, torch.Tensor]]: def extract_normalization_stats(state_dict: dict[str, torch.Tensor]) -> dict[str, dict[str, torch.Tensor]]:
""" """Scans a model's state_dict to find and extract normalization statistics.
Scans a model's state_dict to find and extract normalization statistics.
This function identifies keys corresponding to normalization layers (e.g., those This function identifies keys corresponding to normalization layers (e.g., those
for mean, std, min, max) based on a set of predefined patterns and organizes for mean, std, min, max) based on a set of predefined patterns and organizes
them into a nested dictionary. them into a nested dictionary.
Args: Args:
state_dict: The state dictionary of a pretrained policy model. state_dict (`dict[str, torch.Tensor]`): The model's state dictionary to scan.
Returns: Returns:
A nested dictionary where outer keys are feature names (e.g., A nested dictionary where outer keys are feature names (e.g.,
@@ -125,8 +122,7 @@ def extract_normalization_stats(state_dict: dict[str, torch.Tensor]) -> dict[str
def detect_features_and_norm_modes( def detect_features_and_norm_modes(
config: dict[str, Any], stats: dict[str, dict[str, torch.Tensor]] config: dict[str, Any], stats: dict[str, dict[str, torch.Tensor]]
) -> tuple[dict[str, PolicyFeature], dict[FeatureType, NormalizationMode]]: ) -> tuple[dict[str, PolicyFeature], dict[FeatureType, NormalizationMode]]:
""" """Infers policy features and normalization modes from the model config and stats.
Infers policy features and normalization modes from the model config and stats.
This function first attempts to find feature definitions and normalization This function first attempts to find feature definitions and normalization
mappings directly from the policy's configuration file. If this information is mappings directly from the policy's configuration file. If this information is
@@ -136,8 +132,9 @@ def detect_features_and_norm_modes(
It applies sensible defaults if inference is not possible. It applies sensible defaults if inference is not possible.
Args: Args:
config: The policy's configuration dictionary from `config.json`. config (`dict[str, Any]`): The policy's configuration dictionary (from `config.json`).
stats: The normalization statistics extracted from the model's state_dict. stats (`dict[str, dict[str, torch.Tensor]]`): The normalization statistics extracted by
`extract_normalization_stats`.
Returns: Returns:
A tuple containing: A tuple containing:
@@ -248,14 +245,13 @@ def detect_features_and_norm_modes(
def remove_normalization_layers(state_dict: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: def remove_normalization_layers(state_dict: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
""" """Creates a new state_dict with all normalization-related layers removed.
Creates a new state_dict with all normalization-related layers removed.
This function filters the original state dictionary, excluding any keys that This function filters the original state dictionary, excluding any keys that
match a set of predefined patterns associated with normalization modules. match a set of predefined patterns associated with normalization modules.
Args: Args:
state_dict: The original model state dictionary. state_dict (`dict[str, torch.Tensor]`): The original model state dictionary.
Returns: Returns:
A new state dictionary containing only the core model weights, without A new state dictionary containing only the core model weights, without
@@ -286,12 +282,11 @@ def remove_normalization_layers(state_dict: dict[str, torch.Tensor]) -> dict[str
def clean_state_dict( def clean_state_dict(
state_dict: dict[str, torch.Tensor], remove_str: str = "._orig_mod" state_dict: dict[str, torch.Tensor], remove_str: str = "._orig_mod"
) -> dict[str, torch.Tensor]: ) -> dict[str, torch.Tensor]:
""" """Remove a substring (e.g. '._orig_mod') from all keys in a state dict.
Remove a substring (e.g. '._orig_mod') from all keys in a state dict.
Args: Args:
state_dict (dict): The original state dict. state_dict (dict): The original state dict.
remove_str (str): The substring to remove from the keys. remove_str (str, *optional*, defaults to `"._orig_mod"`): The substring to remove from the keys.
Returns: Returns:
dict: A new state dict with cleaned keys. dict: A new state dict with cleaned keys.
@@ -309,18 +304,18 @@ def load_state_dict_with_missing_key_handling(
policy_type: str, policy_type: str,
known_missing_keys_whitelist: dict[str, list[str]], known_missing_keys_whitelist: dict[str, list[str]],
) -> list[str]: ) -> list[str]:
""" """Load state dict into policy with graceful handling of missing keys.
Load state dict into policy with graceful handling of missing keys.
This function loads the state dict with strict=False, filters out whitelisted This function loads the state dict with strict=False, filters out whitelisted
missing keys, and provides detailed reporting about any issues found. missing keys, and provides detailed reporting about any issues found.
Args: Args:
policy: The policy model to load the state dict into. policy (`torch.nn.Module`): The policy module to load the state dict into.
state_dict: The cleaned state dictionary to load. state_dict (`dict[str, torch.Tensor]`): The cleaned state dict to load.
policy_type: The type of policy (used for whitelist lookup). policy_type (`str`): The policy type name, used to look up the whitelist (matched
known_missing_keys_whitelist: Dictionary mapping policy types to lists of case-insensitively).
known acceptable missing keys. known_missing_keys_whitelist (`dict[str, list[str]]`): A mapping from policy type to the list of
key names that are expected to be missing for that policy.
Returns: Returns:
List of problematic missing keys that weren't in the whitelist. List of problematic missing keys that weren't in the whitelist.
@@ -363,12 +358,11 @@ def load_state_dict_with_missing_key_handling(
def convert_features_to_policy_features(features_dict: dict[str, dict]) -> dict[str, PolicyFeature]: def convert_features_to_policy_features(features_dict: dict[str, dict]) -> dict[str, PolicyFeature]:
""" """Converts a feature dictionary from the old config format to the new `PolicyFeature` format.
Converts a feature dictionary from the old config format to the new `PolicyFeature` format.
Args: Args:
features_dict: The feature dictionary in the old format, where values are features_dict (`dict[str, dict]`): A mapping from feature name to its old-format config dict
simple dictionaries (e.g., `{"shape": [7]}`). (with a `"shape"` or `"dim"` key).
Returns: Returns:
A dictionary mapping feature names to `PolicyFeature` dataclass objects. A dictionary mapping feature names to `PolicyFeature` dataclass objects.
@@ -396,11 +390,10 @@ def convert_features_to_policy_features(features_dict: dict[str, dict]) -> dict[
def display_migration_summary_with_warnings(problematic_missing_keys: list[str]) -> None: def display_migration_summary_with_warnings(problematic_missing_keys: list[str]) -> None:
""" """Display final migration summary with warnings about problematic missing keys.
Display final migration summary with warnings about problematic missing keys.
Args: Args:
problematic_missing_keys: List of missing keys that weren't in the whitelist. problematic_missing_keys (`list[str]`): List of missing keys that weren't in the whitelist.
""" """
if not problematic_missing_keys: if not problematic_missing_keys:
return return
@@ -434,12 +427,11 @@ def display_migration_summary_with_warnings(problematic_missing_keys: list[str])
def load_model_from_hub( def load_model_from_hub(
repo_id: str, revision: str | None = None repo_id: str, revision: str | None = None
) -> tuple[dict[str, torch.Tensor], dict[str, Any], dict[str, Any] | None]: ) -> tuple[dict[str, torch.Tensor], dict[str, Any], dict[str, Any] | None]:
""" """Downloads and loads a model's state_dict and configs from the Hugging Face Hub.
Downloads and loads a model's state_dict and configs from the Hugging Face Hub.
Args: Args:
repo_id: The repository ID on the Hub (e.g., 'lerobot/aloha'). repo_id (`str`): The Hugging Face Hub repo ID of the pretrained model.
revision: The specific git revision (branch, tag, or commit hash) to use. revision (`str | None`, *optional*): The Hub revision (branch, tag, or commit hash) to download.
Returns: Returns:
A tuple containing the model's state dictionary, the policy configuration, A tuple containing the model's state dictionary, the policy configuration,
@@ -470,6 +462,7 @@ def load_model_from_hub(
def main(): def main():
"""CLI entry point: parse arguments and migrate a pretrained policy to the processor pipeline format."""
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
description="Migrate policy models with normalization layers to new pipeline system" description="Migrate policy models with normalization layers to new pipeline system"
) )
@@ -23,8 +23,7 @@ from .pipeline import ComplementaryDataProcessorStep, ProcessorStepRegistry
# with serialized processor configs that reference this name. # with serialized processor configs that reference this name.
@ProcessorStepRegistry.register(name="smolvla_new_line_processor") @ProcessorStepRegistry.register(name="smolvla_new_line_processor")
class NewLineTaskProcessorStep(ComplementaryDataProcessorStep): class NewLineTaskProcessorStep(ComplementaryDataProcessorStep):
""" """A processor step that ensures the 'task' description ends with a newline character.
A processor step that ensures the 'task' description ends with a newline character.
This step is necessary for certain tokenizers (e.g., PaliGemma) that expect a This step is necessary for certain tokenizers (e.g., PaliGemma) that expect a
newline at the end of the prompt. It handles both single string tasks and lists newline at the end of the prompt. It handles both single string tasks and lists
@@ -32,6 +31,7 @@ class NewLineTaskProcessorStep(ComplementaryDataProcessorStep):
""" """
def complementary_data(self, complementary_data): def complementary_data(self, complementary_data):
"""Append a trailing newline to the `"task"` entry, if present, leaving other keys untouched."""
if "task" not in complementary_data: if "task" not in complementary_data:
return complementary_data return complementary_data
@@ -56,4 +56,5 @@ class NewLineTaskProcessorStep(ComplementaryDataProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. Only complementary data is touched; features are unchanged."""
return features return features
+29 -30
View File
@@ -38,8 +38,7 @@ from .pipeline import PolicyProcessorPipeline, ProcessorStep, ProcessorStepRegis
@dataclass @dataclass
class _NormalizationMixin: class _NormalizationMixin:
""" """A mixin class providing core functionality for normalization and unnormalization.
A mixin class providing core functionality for normalization and unnormalization.
This class manages normalization statistics (`stats`), converts them to tensors for This class manages normalization statistics (`stats`), converts them to tensors for
efficient computation, handles device placement, and implements the logic for efficient computation, handles device placement, and implements the logic for
@@ -102,8 +101,7 @@ class _NormalizationMixin:
_stats_explicitly_provided: bool = field(default=False, init=False, repr=False) _stats_explicitly_provided: bool = field(default=False, init=False, repr=False)
def __post_init__(self): def __post_init__(self):
""" """Initializes the mixin after dataclass construction.
Initializes the mixin after dataclass construction.
This method handles the robust deserialization of `features` and `norm_map` This method handles the robust deserialization of `features` and `norm_map`
from JSON-compatible formats (where enums become strings and tuples become from JSON-compatible formats (where enums become strings and tuples become
@@ -157,11 +155,11 @@ class _NormalizationMixin:
def to( def to(
self, device: torch.device | str | None = None, dtype: torch.dtype | None = None self, device: torch.device | str | None = None, dtype: torch.dtype | None = None
) -> _NormalizationMixin: ) -> _NormalizationMixin:
""" """Moves the processor's normalization stats to the specified device and/or dtype.
Moves the processor's normalization stats to the specified device.
Args: Args:
device: The target PyTorch device. device: The target PyTorch device.
dtype: The target floating-point dtype for the stats tensors.
Returns: Returns:
The instance of the class, allowing for method chaining. The instance of the class, allowing for method chaining.
@@ -175,8 +173,7 @@ class _NormalizationMixin:
return self return self
def state_dict(self) -> dict[str, Tensor]: def state_dict(self) -> dict[str, Tensor]:
""" """Returns the normalization statistics as a flat state dictionary.
Returns the normalization statistics as a flat state dictionary.
All tensors are moved to the CPU before being returned, which is standard practice All tensors are moved to the CPU before being returned, which is standard practice
for saving state dictionaries. for saving state dictionaries.
@@ -192,8 +189,7 @@ class _NormalizationMixin:
return flat return flat
def load_state_dict(self, state: dict[str, Tensor]) -> None: def load_state_dict(self, state: dict[str, Tensor]) -> None:
""" """Loads normalization statistics from a state dictionary.
Loads normalization statistics from a state dictionary.
The loaded tensors are moved to the processor's configured device. The loaded tensors are moved to the processor's configured device.
@@ -244,8 +240,7 @@ class _NormalizationMixin:
self.stats[key][stat_name] = from_tensor_to_numpy(tensor) self.stats[key][stat_name] = from_tensor_to_numpy(tensor)
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
""" """Returns a serializable dictionary of the processor's configuration.
Returns a serializable dictionary of the processor's configuration.
This method is used when saving the processor to disk, ensuring that its This method is used when saving the processor to disk, ensuring that its
configuration can be reconstructed later. configuration can be reconstructed later.
@@ -265,8 +260,7 @@ class _NormalizationMixin:
return config return config
def _normalize_observation(self, observation: RobotObservation, inverse: bool) -> dict[str, Tensor]: def _normalize_observation(self, observation: RobotObservation, inverse: bool) -> dict[str, Tensor]:
""" """Applies (un)normalization to all relevant features in an observation dictionary.
Applies (un)normalization to all relevant features in an observation dictionary.
Args: Args:
observation: The observation dictionary to process. observation: The observation dictionary to process.
@@ -287,8 +281,7 @@ class _NormalizationMixin:
def _normalize_action(self, action: Tensor, inverse: bool) -> Tensor: def _normalize_action(self, action: Tensor, inverse: bool) -> Tensor:
# Convert to tensor but preserve original dtype for adaptation logic # Convert to tensor but preserve original dtype for adaptation logic
""" """Applies (un)normalization to an action tensor.
Applies (un)normalization to an action tensor.
Args: Args:
action: The action tensor to process. action: The action tensor to process.
@@ -303,8 +296,7 @@ class _NormalizationMixin:
def _apply_transform( def _apply_transform(
self, tensor: Tensor, key: str, feature_type: FeatureType, *, inverse: bool = False self, tensor: Tensor, key: str, feature_type: FeatureType, *, inverse: bool = False
) -> Tensor: ) -> Tensor:
""" """Core logic to apply a normalization or unnormalization transformation to a tensor.
Core logic to apply a normalization or unnormalization transformation to a tensor.
This method selects the appropriate normalization mode based on the feature type This method selects the appropriate normalization mode based on the feature type
and applies the corresponding mathematical operation. and applies the corresponding mathematical operation.
@@ -425,8 +417,7 @@ class _NormalizationMixin:
@dataclass @dataclass
@ProcessorStepRegistry.register(name="normalizer_processor") @ProcessorStepRegistry.register(name="normalizer_processor")
class NormalizerProcessorStep(_NormalizationMixin, ProcessorStep): class NormalizerProcessorStep(_NormalizationMixin, ProcessorStep):
""" """A processor step that applies normalization to observations and actions in a transition.
A processor step that applies normalization to observations and actions in a transition.
This class uses the logic from `_NormalizationMixin` to perform forward normalization This class uses the logic from `_NormalizationMixin` to perform forward normalization
(e.g., scaling data to have zero mean and unit variance, or to the range [-1, 1]). (e.g., scaling data to have zero mean and unit variance, or to the range [-1, 1]).
@@ -444,8 +435,7 @@ class NormalizerProcessorStep(_NormalizationMixin, ProcessorStep):
eps: float = 1e-8, eps: float = 1e-8,
device: torch.device | str | None = None, device: torch.device | str | None = None,
) -> NormalizerProcessorStep: ) -> NormalizerProcessorStep:
""" """Creates a `NormalizerProcessorStep` instance using statistics from a `LeRobotDataset`.
Creates a `NormalizerProcessorStep` instance using statistics from a `LeRobotDataset`.
Args: Args:
dataset: The dataset from which to extract normalization statistics. dataset: The dataset from which to extract normalization statistics.
@@ -468,6 +458,11 @@ class NormalizerProcessorStep(_NormalizationMixin, ProcessorStep):
) )
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
"""Normalize the transition's observation and action in place (a copy of the transition).
Raises:
ValueError: If the transition has an action that is not a `PolicyAction`.
"""
new_transition = transition.copy() new_transition = transition.copy()
# Handle observation normalization. # Handle observation normalization.
@@ -493,14 +488,14 @@ class NormalizerProcessorStep(_NormalizationMixin, ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. A value transformation; features are unchanged."""
return features return features
@dataclass @dataclass
@ProcessorStepRegistry.register(name="unnormalizer_processor") @ProcessorStepRegistry.register(name="unnormalizer_processor")
class UnnormalizerProcessorStep(_NormalizationMixin, ProcessorStep): class UnnormalizerProcessorStep(_NormalizationMixin, ProcessorStep):
""" """A processor step that applies unnormalization to observations and actions.
A processor step that applies unnormalization to observations and actions.
This class inverts the normalization process, scaling data back to its original This class inverts the normalization process, scaling data back to its original
range. It is typically used in the post-processing pipeline to convert a policy's range. It is typically used in the post-processing pipeline to convert a policy's
@@ -517,8 +512,7 @@ class UnnormalizerProcessorStep(_NormalizationMixin, ProcessorStep):
*, *,
device: torch.device | str | None = None, device: torch.device | str | None = None,
) -> UnnormalizerProcessorStep: ) -> UnnormalizerProcessorStep:
""" """Creates an `UnnormalizerProcessorStep` using statistics from a `LeRobotDataset`.
Creates an `UnnormalizerProcessorStep` using statistics from a `LeRobotDataset`.
Args: Args:
dataset: The dataset from which to extract normalization statistics. dataset: The dataset from which to extract normalization statistics.
@@ -532,6 +526,11 @@ class UnnormalizerProcessorStep(_NormalizationMixin, ProcessorStep):
return cls(features=features, norm_map=norm_map, stats=dataset.meta.stats, device=device) return cls(features=features, norm_map=norm_map, stats=dataset.meta.stats, device=device)
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
"""Unnormalize the transition's observation and action in place (a copy of the transition).
Raises:
ValueError: If the transition has an action that is not a `PolicyAction`.
"""
new_transition = transition.copy() new_transition = transition.copy()
# Handle observation unnormalization. # Handle observation unnormalization.
@@ -554,14 +553,14 @@ class UnnormalizerProcessorStep(_NormalizationMixin, ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. A value transformation; features are unchanged."""
return features return features
def hotswap_stats( def hotswap_stats(
policy_processor: PolicyProcessorPipeline, stats: dict[str, dict[str, Any]] policy_processor: PolicyProcessorPipeline, stats: dict[str, dict[str, Any]]
) -> PolicyProcessorPipeline: ) -> PolicyProcessorPipeline:
""" """Replaces normalization statistics in an existing `PolicyProcessorPipeline` instance.
Replaces normalization statistics in an existing `PolicyProcessorPipeline` instance.
This function creates a deep copy of the provided pipeline and updates the This function creates a deep copy of the provided pipeline and updates the
statistics of any `NormalizerProcessorStep` or `UnnormalizerProcessorStep` it statistics of any `NormalizerProcessorStep` or `UnnormalizerProcessorStep` it
@@ -570,8 +569,8 @@ def hotswap_stats(
pipeline. pipeline.
Args: Args:
policy_processor: The policy processor pipeline to modify. policy_processor (`PolicyProcessorPipeline`): The policy processor pipeline to modify.
stats: The new dictionary of normalization statistics to apply. stats (`dict[str, dict[str, Any]]`): The new dictionary of normalization statistics to apply.
Returns: Returns:
A new `PolicyProcessorPipeline` instance with the updated statistics. A new `PolicyProcessorPipeline` instance with the updated statistics.
+5 -10
View File
@@ -29,8 +29,7 @@ from .pipeline import ObservationProcessorStep, ProcessorStepRegistry
@dataclass @dataclass
@ProcessorStepRegistry.register(name="observation_processor") @ProcessorStepRegistry.register(name="observation_processor")
class VanillaObservationProcessorStep(ObservationProcessorStep): class VanillaObservationProcessorStep(ObservationProcessorStep):
""" """Processes standard Gymnasium observations into the LeRobot format.
Processes standard Gymnasium observations into the LeRobot format.
This step handles both image and state data from a typical observation dictionary, This step handles both image and state data from a typical observation dictionary,
preparing it for use in a LeRobot policy. preparing it for use in a LeRobot policy.
@@ -53,8 +52,7 @@ class VanillaObservationProcessorStep(ObservationProcessorStep):
""" """
def _process_single_image(self, img: np.ndarray) -> Tensor: def _process_single_image(self, img: np.ndarray) -> Tensor:
""" """Processes a single NumPy image array into a channel-first, normalized tensor.
Processes a single NumPy image array into a channel-first, normalized tensor.
Args: Args:
img: A NumPy array representing the image, expected to be in channel-last img: A NumPy array representing the image, expected to be in channel-last
@@ -92,10 +90,7 @@ class VanillaObservationProcessorStep(ObservationProcessorStep):
return img_tensor return img_tensor
def _process_observation(self, observation): def _process_observation(self, observation):
""" """Processes both image and state observations."""
Processes both image and state observations.
"""
processed_obs = observation.copy() processed_obs = observation.copy()
if "pixels" in processed_obs: if "pixels" in processed_obs:
@@ -126,13 +121,13 @@ class VanillaObservationProcessorStep(ObservationProcessorStep):
return processed_obs return processed_obs
def observation(self, observation): def observation(self, observation):
"""See [`~processor.ObservationProcessorStep.observation`]. Delegates to `_process_observation`."""
return self._process_observation(observation) return self._process_observation(observation)
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Transforms feature keys from the Gym standard to the LeRobot standard.
Transforms feature keys from the Gym standard to the LeRobot standard.
This method standardizes the feature dictionary by renaming keys according This method standardizes the feature dictionary by renaming keys according
to LeRobot's conventions, ensuring that policies can be constructed correctly. to LeRobot's conventions, ensuring that policies can be constructed correctly.
+12 -5
View File
@@ -14,11 +14,10 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
""" """This module defines a generic, sequential data processing pipeline framework.
This module defines a generic, sequential data processing pipeline framework, primarily designed for
transforming robotics data (observations, actions, rewards, etc.).
The core components are: It is primarily designed for transforming robotics data (observations, actions, rewards, etc.). The core
components are:
- ProcessorStep: An abstract base class for a single data transformation operation. - ProcessorStep: An abstract base class for a single data transformation operation.
- ProcessorStepRegistry: A mechanism to register and retrieve ProcessorStep classes by name. - ProcessorStepRegistry: A mechanism to register and retrieve ProcessorStep classes by name.
- DataProcessorPipeline: A class that chains multiple ProcessorStep instances together to form a complete - DataProcessorPipeline: A class that chains multiple ProcessorStep instances together to form a complete
@@ -249,9 +248,16 @@ class ProcessorKwargs(TypedDict, total=False):
class ProcessorMigrationError(Exception): class ProcessorMigrationError(Exception):
"""Raised when a model needs migration to the processor format""" """Raised when a model needs migration to the processor format."""
def __init__(self, model_path: str | Path, migration_command: str, original_error: str): def __init__(self, model_path: str | Path, migration_command: str, original_error: str):
"""Build the error message pointing the user at the migration command to run.
Args:
model_path: Path or Hub repo ID of the model that needs migration.
migration_command: Shell command the user should run to migrate it.
original_error: The underlying error that triggered this migration check.
"""
self.model_path = model_path self.model_path = model_path
self.migration_command = migration_command self.migration_command = migration_command
self.original_error = original_error self.original_error = original_error
@@ -1486,6 +1492,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
feature_types = {feature_type.value for feature_type in FeatureType} feature_types = {feature_type.value for feature_type in FeatureType}
def is_policy_feature_mapping(features: Any) -> bool: def is_policy_feature_mapping(features: Any) -> bool:
"""Return `True` if `features` looks like a serialized `dict[str, PolicyFeature]`."""
return ( return (
isinstance(features, dict) isinstance(features, dict)
and bool(features) and bool(features)
@@ -34,14 +34,21 @@ class RobotActionToPolicyActionProcessorStep(ActionProcessorStep):
motor_names: list[str] motor_names: list[str]
def action(self, action: RobotAction) -> PolicyAction: def action(self, action: RobotAction) -> PolicyAction:
"""Stack `action`'s `"{motor}.pos"` entries, in `motor_names` order, into a single tensor.
Raises:
ValueError: If `action` doesn't have exactly `len(motor_names)` entries.
"""
if len(self.motor_names) != len(action): if len(self.motor_names) != len(action):
raise ValueError(f"Action must have {len(self.motor_names)} elements, got {len(action)}") raise ValueError(f"Action must have {len(self.motor_names)} elements, got {len(action)}")
return torch.tensor([action[f"{name}.pos"] for name in self.motor_names]) return torch.tensor([action[f"{name}.pos"] for name in self.motor_names])
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
"""Returns `{"motor_names": [...]}`."""
return asdict(self) return asdict(self)
def transform_features(self, features): def transform_features(self, features):
"""Replace the per-motor action features with a single stacked action feature."""
features[PipelineFeatureType.ACTION][ACTION] = PolicyFeature( features[PipelineFeatureType.ACTION][ACTION] = PolicyFeature(
type=FeatureType.ACTION, shape=(len(self.motor_names),) type=FeatureType.ACTION, shape=(len(self.motor_names),)
) )
@@ -56,14 +63,21 @@ class PolicyActionToRobotActionProcessorStep(ActionProcessorStep):
motor_names: list[str] motor_names: list[str]
def action(self, action: PolicyAction) -> RobotAction: def action(self, action: PolicyAction) -> RobotAction:
"""Split `action`, in `motor_names` order, into a `"{motor}.pos"`-keyed dict.
Raises:
ValueError: If `action` doesn't have exactly `len(motor_names)` elements.
"""
if len(self.motor_names) != len(action): if len(self.motor_names) != len(action):
raise ValueError(f"Action must have {len(self.motor_names)} elements, got {len(action)}") raise ValueError(f"Action must have {len(self.motor_names)} elements, got {len(action)}")
return {f"{name}.pos": action[i] for i, name in enumerate(self.motor_names)} return {f"{name}.pos": action[i] for i, name in enumerate(self.motor_names)}
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
"""Returns `{"motor_names": [...]}`."""
return asdict(self) return asdict(self)
def transform_features(self, features): def transform_features(self, features):
"""Replace the stacked action feature with one per-motor `"{motor}.pos"` action feature."""
for name in self.motor_names: for name in self.motor_names:
features[PipelineFeatureType.ACTION][f"{name}.pos"] = PolicyFeature( features[PipelineFeatureType.ACTION][f"{name}.pos"] = PolicyFeature(
type=FeatureType.ACTION, shape=(1,) type=FeatureType.ACTION, shape=(1,)
@@ -41,9 +41,9 @@ def to_relative_actions(actions: Tensor, state: Tensor, mask: Sequence[bool]) ->
"""Convert absolute actions to relative: relative = action - state (for masked dims). """Convert absolute actions to relative: relative = action - state (for masked dims).
Args: Args:
actions: (B, T, action_dim) or (B, action_dim). actions (`Tensor`): `(B, T, action_dim)` or `(B, action_dim)`.
state: (B, state_dim). Broadcast across time dimension. state (`Tensor`): `(B, state_dim)`. Broadcast across the time dimension.
mask: Which dims to convert. Can be shorter than action_dim. mask (`Sequence[bool]`): Which dims to convert. Can be shorter than `action_dim`.
""" """
mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device) mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device)
dims = mask_t.shape[0] dims = mask_t.shape[0]
@@ -63,9 +63,9 @@ def to_absolute_actions(actions: Tensor, state: Tensor, mask: Sequence[bool]) ->
"""Convert relative actions back to absolute: absolute = relative + state (for masked dims). """Convert relative actions back to absolute: absolute = relative + state (for masked dims).
Args: Args:
actions: (B, T, action_dim) or (B, action_dim). actions (`Tensor`): `(B, T, action_dim)` or `(B, action_dim)`.
state: (B, state_dim). Broadcast across time dimension. state (`Tensor`): `(B, state_dim)`. Broadcast across the time dimension.
mask: Which dims to convert. Can be shorter than action_dim. mask (`Sequence[bool]`): Which dims to convert. Can be shorter than `action_dim`.
""" """
mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device) mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device)
dims = mask_t.shape[0] dims = mask_t.shape[0]
@@ -123,6 +123,7 @@ class RelativeActionsProcessorStep(ProcessorStep):
return mask return mask
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
"""Cache `observation.state` for the paired postprocessing step, and convert `action` to relative if `enabled`."""
observation = transition.get(TransitionKey.OBSERVATION, {}) observation = transition.get(TransitionKey.OBSERVATION, {})
state = observation.get(OBS_STATE) if observation else None state = observation.get(OBS_STATE) if observation else None
@@ -147,6 +148,7 @@ class RelativeActionsProcessorStep(ProcessorStep):
return self._last_state return self._last_state
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
"""Returns `{"enabled": ..., "exclude_joints": ..., "action_names": ...}`."""
return { return {
"enabled": self.enabled, "enabled": self.enabled,
"exclude_joints": self.exclude_joints, "exclude_joints": self.exclude_joints,
@@ -156,6 +158,7 @@ class RelativeActionsProcessorStep(ProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. A value transformation; features are unchanged."""
return features return features
@@ -178,6 +181,11 @@ class AbsoluteActionsProcessorStep(ProcessorStep):
relative_step: RelativeActionsProcessorStep | None = field(default=None, repr=False) relative_step: RelativeActionsProcessorStep | None = field(default=None, repr=False)
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
"""Convert `action` back to absolute using the paired step's cached state, if `enabled`.
Raises:
RuntimeError: If `relative_step` is unset, or no state has been cached yet.
"""
if not self.enabled: if not self.enabled:
return transition return transition
@@ -204,9 +212,11 @@ class AbsoluteActionsProcessorStep(ProcessorStep):
return new_transition return new_transition
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
"""Returns `{"enabled": ...}`."""
return {"enabled": self.enabled} return {"enabled": self.enabled}
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""See [`~processor.ProcessorStep.transform_features`]. A value transformation; features are unchanged."""
return features return features
+9 -8
View File
@@ -25,8 +25,7 @@ from .pipeline import ObservationProcessorStep, ProcessorStepRegistry
@dataclass @dataclass
@ProcessorStepRegistry.register(name="rename_observations_processor") @ProcessorStepRegistry.register(name="rename_observations_processor")
class RenameObservationsProcessorStep(ObservationProcessorStep): class RenameObservationsProcessorStep(ObservationProcessorStep):
""" """A processor step that renames keys in an observation dictionary.
A processor step that renames keys in an observation dictionary.
This step is useful for creating a standardized data interface by mapping keys This step is useful for creating a standardized data interface by mapping keys
from an environment's format to the format expected by a LeRobot policy or from an environment's format to the format expected by a LeRobot policy or
@@ -40,6 +39,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):
"""Rename each key present in `rename_map`; keys not in `rename_map` are kept as-is."""
processed_obs = {} processed_obs = {}
for key, value in observation.items(): for key, value in observation.items():
if key in self.rename_map: if key in self.rename_map:
@@ -50,12 +50,14 @@ class RenameObservationsProcessorStep(ObservationProcessorStep):
return processed_obs return processed_obs
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
"""Returns `{"rename_map": ...}`."""
return {"rename_map": self.rename_map} return {"rename_map": self.rename_map}
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
"""Transforms: """Rename observation feature keys the same way `observation` renames observation data.
- Each key in the observation that appears in `rename_map` is renamed to its value. - Each key in the observation that appears in `rename_map` is renamed to its value.
- Keys not in `rename_map` remain unchanged. - Keys not in `rename_map` remain unchanged.
""" """
@@ -67,17 +69,16 @@ class RenameObservationsProcessorStep(ObservationProcessorStep):
def rename_stats(stats: dict[str, dict[str, Any]], rename_map: dict[str, str]) -> dict[str, dict[str, Any]]: def rename_stats(stats: dict[str, dict[str, Any]], rename_map: dict[str, str]) -> dict[str, dict[str, Any]]:
""" """Renames the top-level keys in a statistics dictionary using a provided mapping.
Renames the top-level keys in a statistics dictionary using a provided mapping.
This is a helper function typically used to keep normalization statistics This is a helper function typically used to keep normalization statistics
consistent with renamed observation or action features. It performs a defensive consistent with renamed observation or action features. It performs a defensive
deep copy to avoid modifying the original `stats` dictionary. deep copy to avoid modifying the original `stats` dictionary.
Args: Args:
stats: A nested dictionary of statistics, where top-level keys are stats (`dict[str, dict[str, Any]]`): A nested dictionary of statistics, where top-level keys are
feature names (e.g., `{"observation.state": {"mean": 0.5}}`). feature names (e.g., `{"observation.state": {"mean": 0.5}}`).
rename_map: A dictionary mapping old feature names to new feature names. rename_map (`dict[str, str]`): A dictionary mapping old feature names to new feature names.
Returns: Returns:
A new statistics dictionary with its top-level keys renamed. Returns an A new statistics dictionary with its top-level keys renamed. Returns an
@@ -48,10 +48,12 @@ class RenderMessagesStep(ProcessorStep):
dataset_ctx: Any | None = None dataset_ctx: Any | None = None
def __post_init__(self) -> None: def __post_init__(self) -> None:
"""Deserialize `recipe` from a plain dict, if it was passed as one (e.g. loaded from JSON config)."""
if isinstance(self.recipe, dict): if isinstance(self.recipe, dict):
self.recipe = TrainingRecipe.from_dict(self.recipe) self.recipe = TrainingRecipe.from_dict(self.recipe)
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
"""Returns `{"recipe": ...}`, with `recipe` serialized to a plain dict."""
return {"recipe": asdict(self.recipe)} return {"recipe": asdict(self.recipe)}
def __call__(self, transition: EnvTransition) -> EnvTransition | None: def __call__(self, transition: EnvTransition) -> EnvTransition | None:
+25 -38
View File
@@ -14,8 +14,7 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
""" """This script defines a processor for tokenizing natural language instructions from an environment transition.
This script defines a processor for tokenizing natural language instructions from an environment transition.
It uses a tokenizer from the Hugging Face `transformers` library to convert task descriptions (text) into It uses a tokenizer from the Hugging Face `transformers` library to convert task descriptions (text) into
token IDs and attention masks, which are then added to the observation dictionary. token IDs and attention masks, which are then added to the observation dictionary.
@@ -56,8 +55,7 @@ else:
@dataclass @dataclass
@ProcessorStepRegistry.register(name="tokenizer_processor") @ProcessorStepRegistry.register(name="tokenizer_processor")
class TokenizerProcessorStep(ObservationProcessorStep): class TokenizerProcessorStep(ObservationProcessorStep):
""" """Processor step to tokenize a natural language task description.
Processor step to tokenize a natural language task description.
This step extracts a task string from the `complementary_data` of an `EnvTransition`, This step extracts a task string from the `complementary_data` of an `EnvTransition`,
tokenizes it using a Hugging Face `transformers` tokenizer, and adds the resulting tokenizes it using a Hugging Face `transformers` tokenizer, and adds the resulting
@@ -90,8 +88,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
input_tokenizer: Any = field(default=None, init=False, repr=False) input_tokenizer: Any = field(default=None, init=False, repr=False)
def __post_init__(self): def __post_init__(self):
""" """Initializes the tokenizer after the dataclass is created.
Initializes the tokenizer after the dataclass is created.
It checks for the availability of the `transformers` library and loads the tokenizer It checks for the availability of the `transformers` library and loads the tokenizer
either from a provided object or by name from the Hugging Face Hub. either from a provided object or by name from the Hugging Face Hub.
@@ -120,8 +117,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
) )
def get_task(self, transition: EnvTransition) -> list[str] | None: def get_task(self, transition: EnvTransition) -> list[str] | None:
""" """Extracts the task description(s) from the transition's complementary data.
Extracts the task description(s) from the transition's complementary data.
Args: Args:
transition: The environment transition. transition: The environment transition.
@@ -146,8 +142,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
return None return None
def get_subtask(self, transition: EnvTransition) -> list[str] | None: def get_subtask(self, transition: EnvTransition) -> list[str] | None:
""" """Extracts the subtask from the transition's complementary data.
Extracts the subtask from the transition's complementary data.
Args: Args:
transition: The environment transition. transition: The environment transition.
@@ -172,8 +167,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
return None return None
def observation(self, observation: RobotObservation) -> RobotObservation: def observation(self, observation: RobotObservation) -> RobotObservation:
""" """Tokenizes the task description and adds it to the observation dictionary.
Tokenizes the task description and adds it to the observation dictionary.
This method retrieves the task, tokenizes it, moves the resulting tensors to the This method retrieves the task, tokenizes it, moves the resulting tensors to the
same device as other data in the transition, and updates the observation. same device as other data in the transition, and updates the observation.
@@ -229,8 +223,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
return new_observation return new_observation
def _detect_device(self, transition: EnvTransition) -> torch.device | None: def _detect_device(self, transition: EnvTransition) -> torch.device | None:
""" """Detects the torch.device from existing tensors in the transition.
Detects the torch.device from existing tensors in the transition.
It checks tensors in the observation dictionary first, then the action tensor. It checks tensors in the observation dictionary first, then the action tensor.
@@ -255,8 +248,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
return None # No tensors found, default will be CPU return None # No tensors found, default will be CPU
def _tokenize_text(self, text: str | list[str]) -> dict[str, torch.Tensor]: def _tokenize_text(self, text: str | list[str]) -> dict[str, torch.Tensor]:
""" """A wrapper around the tokenizer call.
A wrapper around the tokenizer call.
Args: Args:
text: A string or list of strings to tokenize. text: A string or list of strings to tokenize.
@@ -274,8 +266,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
) )
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
""" """Returns the serializable configuration of the processor.
Returns the serializable configuration of the processor.
Note: The tokenizer object itself is not serialized. If the processor was initialized Note: The tokenizer object itself is not serialized. If the processor was initialized
with a tokenizer name, that name will be included in the config. with a tokenizer name, that name will be included in the config.
@@ -309,8 +300,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Adds feature definitions for the language tokens and attention mask.
Adds feature definitions for the language tokens and attention mask.
This updates the policy features dictionary to include the new data added to the This updates the policy features dictionary to include the new data added to the
observation, ensuring downstream components are aware of their shape and type. observation, ensuring downstream components are aware of their shape and type.
@@ -339,8 +329,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
@dataclass @dataclass
@ProcessorStepRegistry.register(name="action_tokenizer_processor") @ProcessorStepRegistry.register(name="action_tokenizer_processor")
class ActionTokenizerProcessorStep(ActionProcessorStep): class ActionTokenizerProcessorStep(ActionProcessorStep):
""" """Processor step to tokenize action data using a fast action tokenizer.
Processor step to tokenize action data using a fast action tokenizer.
This step takes action tensors from an `EnvTransition`, tokenizes them using This step takes action tensors from an `EnvTransition`, tokenizes them using
a Hugging Face `transformers` AutoProcessor (such as the Physical Intelligence "fast" tokenizer), a Hugging Face `transformers` AutoProcessor (such as the Physical Intelligence "fast" tokenizer),
@@ -373,8 +362,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
_paligemma_tokenizer: Any = field(default=None, init=False, repr=False) _paligemma_tokenizer: Any = field(default=None, init=False, repr=False)
def __post_init__(self): def __post_init__(self):
""" """Initializes the action tokenizer after the dataclass is created.
Initializes the action tokenizer after the dataclass is created.
It checks for the availability of the `transformers` library and loads the tokenizer It checks for the availability of the `transformers` library and loads the tokenizer
either from a provided object or by name from the Hugging Face Hub. either from a provided object or by name from the Hugging Face Hub.
@@ -412,8 +400,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
) )
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
""" """Applies action tokenization to the transition.
Applies action tokenization to the transition.
This overrides the base class to handle both tokens and mask. This overrides the base class to handle both tokens and mask.
@@ -445,14 +432,11 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
return new_transition return new_transition
def _act_tokens_to_paligemma_tokens(self, tokens: torch.Tensor) -> torch.Tensor: def _act_tokens_to_paligemma_tokens(self, tokens: torch.Tensor) -> torch.Tensor:
""" """Converts action tokens to PaliGemma tokens."""
Converts action tokens to PaliGemma tokens.
"""
return self._paligemma_tokenizer.vocab_size - 1 - self.fast_skip_tokens - tokens return self._paligemma_tokenizer.vocab_size - 1 - self.fast_skip_tokens - tokens
def _tokenize_action(self, action: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: def _tokenize_action(self, action: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
""" """Tokenizes the action tensor and creates a mask.
Tokenizes the action tensor and creates a mask.
Args: Args:
action: The input action tensor to tokenize. Shape: (B, H, action_dim) or (H, action_dim,) action: The input action tensor to tokenize. Shape: (B, H, action_dim) or (H, action_dim,)
@@ -568,16 +552,15 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
return tokens_batch, masks_batch, code_masks_batch return tokens_batch, masks_batch, code_masks_batch
def action(self, action: torch.Tensor) -> torch.Tensor: def action(self, action: torch.Tensor) -> torch.Tensor:
""" """This method is not used since we override `__call__`.
This method is not used since we override __call__.
Required by ActionProcessorStep ABC. Required by the `ActionProcessorStep` ABC.
""" """
tokens, _, _ = self._tokenize_action(action) tokens, _, _ = self._tokenize_action(action)
return tokens return tokens
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
""" """Returns the serializable configuration of the processor.
Returns the serializable configuration of the processor.
Note: The tokenizer object itself is not serialized. If the processor was initialized Note: The tokenizer object itself is not serialized. If the processor was initialized
with a tokenizer name, that name will be included in the config. with a tokenizer name, that name will be included in the config.
@@ -600,6 +583,11 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
return config return config
def save_artifacts(self, save_directory: Path) -> dict[str, str]: def save_artifacts(self, save_directory: Path) -> dict[str, str]:
"""Save the action tokenizer so object-provided instances reload without overrides.
Raises:
TypeError: If `action_tokenizer` doesn't implement `save_pretrained`.
"""
artifact_path = Path("action_tokenizer") artifact_path = Path("action_tokenizer")
save_pretrained = getattr(self.action_tokenizer, "save_pretrained", None) save_pretrained = getattr(self.action_tokenizer, "save_pretrained", None)
if save_pretrained is None: if save_pretrained is None:
@@ -610,8 +598,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
def transform_features( def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
""" """Updates feature definitions to reflect tokenized actions.
Updates feature definitions to reflect tokenized actions.
This updates the policy features dictionary to indicate that the action This updates the policy features dictionary to indicate that the action
has been tokenized into a sequence of token IDs with shape (max_action_tokens,). has been tokenized into a sequence of token IDs with shape (max_action_tokens,).
+1
View File
@@ -60,6 +60,7 @@ PATH_TO_LEROBOT = PATH_TO_REPO / "src" / "lerobot"
# Modules whose public objects are checked. Add a module here once its docstrings follow the standard. # Modules whose public objects are checked. Add a module here once its docstrings follow the standard.
MODULES_TO_CHECK = [ MODULES_TO_CHECK = [
"lerobot.robots", "lerobot.robots",
"lerobot.processor",
] ]
# Objects that do not yet follow the standard, so the check can be green from day one. Removing an entry # Objects that do not yet follow the standard, so the check can be green from day one. Removing an entry