mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
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:
@@ -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
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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).
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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,).
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user