diff --git a/docs/source/action_representations.mdx b/docs/source/action_representations.mdx index 1604ed467..9ab23f92c 100644 --- a/docs/source/action_representations.mdx +++ b/docs/source/action_representations.mdx @@ -33,7 +33,7 @@ LeRobot provides processor steps for converting between joint and EE spaces usin ```python from lerobot.model.kinematics import RobotKinematics from lerobot.robots.so_follower.robot_kinematic_processor import ( - ForwardKinematicsJointsToEE, + ForwardKinematicsJointsToEEObservation, InverseKinematicsEEToJoints, ) @@ -44,7 +44,7 @@ kinematics = RobotKinematics( ) # Joints → EE (for observations: "where is my gripper?") -fk_step = ForwardKinematicsJointsToEE(kinematics=kinematics, motor_names=[...]) +fk_step = ForwardKinematicsJointsToEEObservation(kinematics=kinematics, motor_names=[...]) # EE → Joints (for actions: "move my gripper here") ik_step = InverseKinematicsEEToJoints(kinematics=kinematics, motor_names=[...]) @@ -197,7 +197,7 @@ Here is how the different processors compose. Each arrow is a processor step, an ``` ┌─────────────────────────────────────────┐ Action Space │ Joint Space ←──IK──→ EE Space │ - │ ForwardKinematicsJointsToEE │ + │ ForwardKinematicsJointsToEEAction │ │ InverseKinematicsEEToJoints │ └─────────────────────────────────────────┘ diff --git a/docs/source/hilserl.mdx b/docs/source/hilserl.mdx index 09a370f3d..ce9ed1999 100644 --- a/docs/source/hilserl.mdx +++ b/docs/source/hilserl.mdx @@ -145,7 +145,7 @@ The environment processor (`env_processor`) handles incoming observations and en 1. **VanillaObservationProcessorStep**: Converts raw robot observations into standardized format 2. **JointVelocityProcessorStep** (optional): Adds joint velocity information to observations 3. **MotorCurrentProcessorStep** (optional): Adds motor current readings to observations -4. **ForwardKinematicsJointsToEE** (optional): Computes end-effector pose from joint positions +4. **ForwardKinematicsJointsToEEObservation** (optional): Computes end-effector pose from joint positions 5. **ImageCropResizeProcessorStep** (optional): Crops and resizes camera images 6. **TimeLimitProcessorStep** (optional): Enforces episode time limits 7. **GripperPenaltyProcessorStep** (optional): Applies penalties for inappropriate gripper usage @@ -413,7 +413,7 @@ We support using a gamepad or a keyboard or the leader arm of the robot. HIL-Serl learns actions in the end-effector space of the robot. Therefore, the teleoperation will control the end-effector's x,y,z displacements. -The end-effector transformation is applied by the processor pipeline (`InverseKinematicsRLStep`, `EEBoundsAndSafety`, `EEReferenceAndDelta`, `GripperVelocityToJoint`) configured under `env.processor.inverse_kinematics` (`InverseKinematicsConfig`) and `env.processor.gripper` / `env.processor.max_gripper_pos`. The defaults related to the end-effector space are: +The end-effector transformation is applied by the processor pipeline (`EEReferenceAndDelta`, `EEBoundsAndSafety`, `GripperVelocityToJoint`, `InverseKinematicsEEToJoints`, `AddIKSolutionStep`) configured under `env.processor.inverse_kinematics` (`InverseKinematicsConfig`) and `env.processor.gripper` / `env.processor.max_gripper_pos`. The defaults related to the end-effector space are: ```python diff --git a/docs/source/phone_teleop.mdx b/docs/source/phone_teleop.mdx index ae79531ef..419ea2c1e 100644 --- a/docs/source/phone_teleop.mdx +++ b/docs/source/phone_teleop.mdx @@ -187,7 +187,7 @@ We use different IK initial guesses in the kinematic steps. As initial guess eit - EEBoundsAndSafety: clamps the EE pose to a workspace and rate‑limits jumps for safety. Also declares `action.ee.*` features. - InverseKinematicsEEToJoints: turns an EE pose into joint positions with IK. `initial_guess_current_joints=True` is recommended for closed‑loop control; set `False` for open‑loop replay for stability. - GripperVelocityToJoint: integrates a velocity‑like gripper input into an absolute gripper position using the current measured state. -- ForwardKinematicsJointsToEE: computes `observation.state.ee.*` from observed joints for logging and training on EE state. +- ForwardKinematicsJointsToEEObservation: computes `observation.state.ee.*` from observed joints for logging and training on EE state. ### Troubleshooting diff --git a/docs/source/processors_robots_teleop.mdx b/docs/source/processors_robots_teleop.mdx index 359b37aff..3ba668409 100644 --- a/docs/source/processors_robots_teleop.mdx +++ b/docs/source/processors_robots_teleop.mdx @@ -58,7 +58,7 @@ robot_ee_to_joints_processor = RobotProcessorPipeline[RobotAction, RobotAction]( robot_joints_to_ee_pose = RobotProcessorPipeline[RobotObservation, RobotObservation]( # robot obs -> dataset obs steps=[ - ForwardKinematicsJointsToEE(kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys())) + ForwardKinematicsJointsToEEObservation(kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys())) ], to_transition=observation_to_transition, to_output=transition_to_observation, diff --git a/examples/phone_to_so100/evaluate.py b/examples/phone_to_so100/evaluate.py index 03925b1e1..8af3409c7 100644 --- a/examples/phone_to_so100/evaluate.py +++ b/examples/phone_to_so100/evaluate.py @@ -36,7 +36,7 @@ from lerobot.processor import ( ) from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig from lerobot.robots.so_follower.robot_kinematic_processor import ( - ForwardKinematicsJointsToEE, + ForwardKinematicsJointsToEEObservation, InverseKinematicsEEToJoints, ) from lerobot.utils.constants import ACTION, OBS_STR @@ -95,7 +95,7 @@ def main(): # Build pipeline to convert joints observation to EE observation robot_joints_to_ee_pose_processor = RobotProcessorPipeline[RobotObservation, RobotObservation]( steps=[ - ForwardKinematicsJointsToEE( + ForwardKinematicsJointsToEEObservation( kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys()) ) ], diff --git a/examples/phone_to_so100/record.py b/examples/phone_to_so100/record.py index 826bee1a1..079808edb 100644 --- a/examples/phone_to_so100/record.py +++ b/examples/phone_to_so100/record.py @@ -29,7 +29,7 @@ from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig from lerobot.robots.so_follower.robot_kinematic_processor import ( EEBoundsAndSafety, EEReferenceAndDelta, - ForwardKinematicsJointsToEE, + ForwardKinematicsJointsToEEObservation, GripperVelocityToJoint, InverseKinematicsEEToJoints, ) @@ -111,7 +111,7 @@ def main(): # Build pipeline to convert joint observation to EE observation (FK). robot_joints_to_ee_pose = RobotProcessorPipeline[RobotObservation, RobotObservation]( steps=[ - ForwardKinematicsJointsToEE( + ForwardKinematicsJointsToEEObservation( kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys()) ) ], diff --git a/examples/phone_to_so100/rollout.py b/examples/phone_to_so100/rollout.py index c3f014d63..8b655ad7b 100644 --- a/examples/phone_to_so100/rollout.py +++ b/examples/phone_to_so100/rollout.py @@ -38,7 +38,7 @@ from lerobot.processor import ( ) from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig from lerobot.robots.so_follower.robot_kinematic_processor import ( - ForwardKinematicsJointsToEE, + ForwardKinematicsJointsToEEObservation, InverseKinematicsEEToJoints, ) from lerobot.rollout import BaseStrategyConfig, RolloutConfig, build_rollout_context @@ -75,7 +75,7 @@ def main(): ) robot_joints_to_ee_pose_processor = RobotProcessorPipeline[RobotObservation, RobotObservation]( - steps=[ForwardKinematicsJointsToEE(kinematics=kinematics_solver, motor_names=motor_names)], + steps=[ForwardKinematicsJointsToEEObservation(kinematics=kinematics_solver, motor_names=motor_names)], to_transition=observation_to_transition, to_output=transition_to_observation, ) diff --git a/examples/so100_to_so100_EE/evaluate.py b/examples/so100_to_so100_EE/evaluate.py index 48ae5bf11..9e097761e 100644 --- a/examples/so100_to_so100_EE/evaluate.py +++ b/examples/so100_to_so100_EE/evaluate.py @@ -36,7 +36,7 @@ from lerobot.processor import ( ) from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig from lerobot.robots.so_follower.robot_kinematic_processor import ( - ForwardKinematicsJointsToEE, + ForwardKinematicsJointsToEEObservation, InverseKinematicsEEToJoints, ) from lerobot.utils.constants import ACTION, OBS_STR @@ -95,7 +95,7 @@ def main(): # Build pipeline to convert joints observation to EE observation robot_joints_to_ee_pose_processor = RobotProcessorPipeline[RobotObservation, RobotObservation]( steps=[ - ForwardKinematicsJointsToEE( + ForwardKinematicsJointsToEEObservation( kinematics=kinematics_solver, motor_names=list(robot.bus.motors.keys()) ) ], diff --git a/examples/so100_to_so100_EE/record.py b/examples/so100_to_so100_EE/record.py index a8e49bdaf..0e14edc81 100644 --- a/examples/so100_to_so100_EE/record.py +++ b/examples/so100_to_so100_EE/record.py @@ -29,7 +29,8 @@ from lerobot.processor import ( from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig from lerobot.robots.so_follower.robot_kinematic_processor import ( EEBoundsAndSafety, - ForwardKinematicsJointsToEE, + ForwardKinematicsJointsToEEAction, + ForwardKinematicsJointsToEEObservation, InverseKinematicsEEToJoints, ) from lerobot.scripts.lerobot_record import record_loop @@ -78,7 +79,7 @@ def main(): # Build pipeline to convert follower joints to EE observation. follower_joints_to_ee = RobotProcessorPipeline[RobotObservation, RobotObservation]( steps=[ - ForwardKinematicsJointsToEE( + ForwardKinematicsJointsToEEObservation( kinematics=follower_kinematics_solver, motor_names=list(follower.bus.motors.keys()) ), ], @@ -89,7 +90,7 @@ def main(): # Build pipeline to convert leader joints to EE action. leader_joints_to_ee = RobotProcessorPipeline[tuple[RobotAction, RobotObservation], RobotAction]( steps=[ - ForwardKinematicsJointsToEE( + ForwardKinematicsJointsToEEAction( kinematics=leader_kinematics_solver, motor_names=list(leader.bus.motors.keys()) ), ], diff --git a/examples/so100_to_so100_EE/rollout.py b/examples/so100_to_so100_EE/rollout.py index a55db1aad..3e3bfb9d0 100644 --- a/examples/so100_to_so100_EE/rollout.py +++ b/examples/so100_to_so100_EE/rollout.py @@ -36,7 +36,7 @@ from lerobot.processor import ( ) from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig from lerobot.robots.so_follower.robot_kinematic_processor import ( - ForwardKinematicsJointsToEE, + ForwardKinematicsJointsToEEObservation, InverseKinematicsEEToJoints, ) from lerobot.rollout import BaseStrategyConfig, RolloutConfig, build_rollout_context @@ -78,7 +78,7 @@ def main(): # Joint-space observation → EE-space observation (consumed by the policy). robot_joints_to_ee_pose_processor = RobotProcessorPipeline[RobotObservation, RobotObservation]( - steps=[ForwardKinematicsJointsToEE(kinematics=kinematics_solver, motor_names=motor_names)], + steps=[ForwardKinematicsJointsToEEObservation(kinematics=kinematics_solver, motor_names=motor_names)], to_transition=observation_to_transition, to_output=transition_to_observation, ) diff --git a/examples/so100_to_so100_EE/teleoperate.py b/examples/so100_to_so100_EE/teleoperate.py index fcad09621..467cf9b27 100644 --- a/examples/so100_to_so100_EE/teleoperate.py +++ b/examples/so100_to_so100_EE/teleoperate.py @@ -27,7 +27,7 @@ from lerobot.processor import ( from lerobot.robots.so_follower import SO100Follower, SO100FollowerConfig from lerobot.robots.so_follower.robot_kinematic_processor import ( EEBoundsAndSafety, - ForwardKinematicsJointsToEE, + ForwardKinematicsJointsToEEAction, InverseKinematicsEEToJoints, ) from lerobot.teleoperators.so_leader import SO100Leader, SO100LeaderConfig @@ -65,7 +65,7 @@ def main(): # Build pipeline to convert teleop joints to EE action leader_to_ee = RobotProcessorPipeline[RobotAction, RobotAction]( steps=[ - ForwardKinematicsJointsToEE( + ForwardKinematicsJointsToEEAction( kinematics=leader_kinematics_solver, motor_names=list(leader.bus.motors.keys()) ), ], diff --git a/src/lerobot/policies/groot/processor_groot.py b/src/lerobot/policies/groot/processor_groot.py index 3cdd25b18..a58f98083 100644 --- a/src/lerobot/policies/groot/processor_groot.py +++ b/src/lerobot/policies/groot/processor_groot.py @@ -56,6 +56,7 @@ from lerobot.processor import ( AddBatchDimensionProcessorStep, DeviceProcessorStep, PolicyAction, + PolicyActionProcessorStep, PolicyProcessorPipeline, ProcessorStep, ProcessorStepRegistry, @@ -2297,7 +2298,7 @@ def _apply_n1_7_action_decode_transform( @dataclass @ProcessorStepRegistry.register(name="groot_n1_7_action_decode_v1") -class GrootN17ActionDecodeStep(ProcessorStep): +class GrootN17ActionDecodeStep(PolicyActionProcessorStep): """Decode the full 132-D N1.7 model action back to environment actions. N1.7 predicts checkpoint-order action groups. This step unnormalizes each @@ -2318,6 +2319,8 @@ class GrootN17ActionDecodeStep(ProcessorStep): and chunk index alongside each queued action through the postprocessor. """ + skip_if_missing = True + env_action_dim: int = 0 raw_stats: dict[str, Any] | None = None modality_config: dict[str, Any] | None = None @@ -2326,20 +2329,17 @@ class GrootN17ActionDecodeStep(ProcessorStep): action_decode_transform: str | None = None pack_step: GrootN17PackInputsStep | None = field(default=None, repr=False) - def __call__(self, transition: EnvTransition) -> EnvTransition: - action = transition.get(TransitionKey.ACTION) - if not isinstance(action, torch.Tensor): - return transition + def action(self, action: PolicyAction) -> PolicyAction: if self.raw_stats is None or self.modality_config is None: - return transition + return action action_config = self.modality_config.get("action", {}) if not isinstance(action_config, dict): - return transition + return action action_keys = action_config.get("modality_keys", []) action_configs = action_config.get("action_configs", []) if not isinstance(action_keys, list) or not isinstance(action_configs, list): - return transition + return action action_np = action.detach().cpu().float().numpy() if self.use_relative_action and action_np.ndim != 3: @@ -2420,7 +2420,7 @@ class GrootN17ActionDecodeStep(ProcessorStep): raise ValueError(f"Unsupported relative N1.7 action config for '{key}': {cfg}") if not decoded_groups: - return transition + return action decoded = np.concatenate( [decoded_groups[key] for key in action_keys if isinstance(key, str) and key in decoded_groups], @@ -2436,11 +2436,7 @@ class GrootN17ActionDecodeStep(ProcessorStep): ) if squeeze_horizon: decoded = decoded[:, 0] - new_transition = transition.copy() - new_transition[TransitionKey.ACTION] = torch.as_tensor( - decoded, dtype=action.dtype, device=action.device - ) - return new_transition + return torch.as_tensor(decoded, dtype=action.dtype, device=action.device) def transform_features(self, features): return features @@ -2461,7 +2457,9 @@ class GrootN17ActionDecodeStep(ProcessorStep): # silently load into it (v1 is stubbed below with the removal guidance). @dataclass @ProcessorStepRegistry.register(name="groot_action_unpack_unnormalize_v2") -class GrootActionUnpackUnnormalizeStep(ProcessorStep): +class GrootActionUnpackUnnormalizeStep(PolicyActionProcessorStep): + skip_if_missing = True + env_action_dim: int = 0 # Apply inverse of min-max normalization if it was used in preprocessor normalize_min_max: bool = True @@ -2470,12 +2468,8 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep): libero_gripper_action: bool = False libero_gripper_binarize: bool = True - def __call__(self, transition: EnvTransition) -> EnvTransition: - # Expect model outputs to be in TransitionKey.ACTION as (B, T, D_model) - action = transition.get(TransitionKey.ACTION) - if not isinstance(action, torch.Tensor): - return transition - + def action(self, action: PolicyAction) -> PolicyAction: + # Model outputs arrive as (B, T, D_model). # Slice to env dimension while preserving an optional action horizon. # Sync rollout postprocesses selected actions as (B, D); RTC postprocesses # chunks as (B, T, D), matching Isaac-GR00T's decode_action contract. @@ -2517,8 +2511,7 @@ class GrootActionUnpackUnnormalizeStep(ProcessorStep): action = action.clone() action[..., -1] = gripper - transition[TransitionKey.ACTION] = action - return transition + return action def transform_features(self, features): return features diff --git a/src/lerobot/policies/molmoact2/processor_molmoact2.py b/src/lerobot/policies/molmoact2/processor_molmoact2.py index 701092225..9b8934d9d 100644 --- a/src/lerobot/policies/molmoact2/processor_molmoact2.py +++ b/src/lerobot/policies/molmoact2/processor_molmoact2.py @@ -41,7 +41,9 @@ from lerobot.processor import ( AddBatchDimensionProcessorStep, DeviceProcessorStep, NormalizerProcessorStep, + ObservationProcessorStep, PolicyAction, + PolicyActionProcessorStep, PolicyProcessorPipeline, ProcessorStep, ProcessorStepRegistry, @@ -1007,7 +1009,7 @@ class MolmoAct2PackInputsProcessorStep(ProcessorStep): @ProcessorStepRegistry.register(name="molmoact2_state_frame_transform") @dataclass -class MolmoAct2StateFrameTransformStep(ProcessorStep): +class MolmoAct2StateFrameTransformStep(ObservationProcessorStep): """Convert robot state from arm frame to model frame before normalization. Required for zero-shot deployment of MolmoAct2-SO100_101 on SO-100/101 @@ -1023,25 +1025,21 @@ class MolmoAct2StateFrameTransformStep(ProcessorStep): See: https://huggingface.co/docs/lerobot/backwardcomp """ + skip_if_missing = True + joint_signs: list[float] | None = None joint_offsets: list[float] | None = None - def __call__(self, transition: EnvTransition) -> EnvTransition: - if self.joint_signs is None or self.joint_offsets is None: - return transition - observation = transition.get(TransitionKey.OBSERVATION) - if not isinstance(observation, dict) or OBS_STATE not in observation: - return transition - transition = transition.copy() - observation = observation.copy() + def observation(self, observation: dict[str, Any]) -> dict[str, Any]: + if self.joint_signs is None or self.joint_offsets is None or OBS_STATE not in observation: + return observation state = torch.as_tensor(observation[OBS_STATE], dtype=torch.float32).clone() n = len(self.joint_signs) signs = torch.tensor(self.joint_signs, dtype=torch.float32, device=state.device) offsets = torch.tensor(self.joint_offsets, dtype=torch.float32, device=state.device) state[..., :n] = signs * state[..., :n] + offsets observation[OBS_STATE] = state - transition[TransitionKey.OBSERVATION] = observation - return transition + return observation def transform_features( self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] @@ -1054,7 +1052,7 @@ class MolmoAct2StateFrameTransformStep(ProcessorStep): @ProcessorStepRegistry.register(name="molmoact2_action_frame_transform") @dataclass -class MolmoAct2ActionFrameTransformStep(ProcessorStep): +class MolmoAct2ActionFrameTransformStep(PolicyActionProcessorStep): """Convert model action from model frame back to arm frame after unnormalization. Inverse of MolmoAct2StateFrameTransformStep. Required for zero-shot @@ -1065,23 +1063,20 @@ class MolmoAct2ActionFrameTransformStep(ProcessorStep): See: https://huggingface.co/docs/lerobot/backwardcomp """ + skip_if_missing = True + joint_signs: list[float] | None = None joint_offsets: list[float] | None = None - def __call__(self, transition: EnvTransition) -> EnvTransition: + def action(self, action: PolicyAction) -> PolicyAction: if self.joint_signs is None or self.joint_offsets is None: - return transition - action = transition.get(TransitionKey.ACTION) - if action is None: - return transition - transition = transition.copy() + return action action = torch.as_tensor(action, dtype=torch.float32).clone() n = len(self.joint_signs) signs = torch.tensor(self.joint_signs, dtype=torch.float32, device=action.device) offsets = torch.tensor(self.joint_offsets, dtype=torch.float32, device=action.device) action[..., :n] = signs * (action[..., :n] - offsets) - transition[TransitionKey.ACTION] = action - return transition + return action def transform_features( self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] @@ -1094,13 +1089,11 @@ class MolmoAct2ActionFrameTransformStep(ProcessorStep): @ProcessorStepRegistry.register(name="molmoact2_clamp_action") @dataclass -class MolmoAct2ClampActionProcessorStep(ProcessorStep): - def __call__(self, transition: EnvTransition) -> EnvTransition: - transition = transition.copy() - action = transition.get(TransitionKey.ACTION) - if action is not None: - transition[TransitionKey.ACTION] = torch.as_tensor(action).clamp(-1.0, 1.0) - return transition +class MolmoAct2ClampActionProcessorStep(PolicyActionProcessorStep): + skip_if_missing = True + + def action(self, action: PolicyAction) -> PolicyAction: + return action.clamp(-1.0, 1.0) def transform_features( self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] diff --git a/src/lerobot/policies/pi05/processor_pi05.py b/src/lerobot/policies/pi05/processor_pi05.py index a0b5a9f0e..b079587c5 100644 --- a/src/lerobot/policies/pi05/processor_pi05.py +++ b/src/lerobot/policies/pi05/processor_pi05.py @@ -14,7 +14,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -from copy import deepcopy from dataclasses import dataclass from typing import Any @@ -22,9 +21,10 @@ import numpy as np import torch from lerobot.configs import PipelineFeatureType, PolicyFeature -from lerobot.lerobot_types import EnvTransition, TransitionKey +from lerobot.lerobot_types import TransitionKey from lerobot.processor import ( AbsoluteActionsProcessorStep, + ComplementaryDataProcessorStep, PolicyAction, PolicyProcessorPipeline, ProcessorStep, @@ -41,7 +41,7 @@ from .configuration_pi05 import PI05Config @ProcessorStepRegistry.register(name="pi05_prepare_state_tokenizer_processor_step") @dataclass -class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep): +class Pi05PrepareStateTokenizerProcessorStep(ComplementaryDataProcessorStep): """ Processor step to prepare the state and tokenize the language input. """ @@ -49,19 +49,14 @@ class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep): max_state_dim: int = 32 task_key: str = "task" - def __call__(self, transition: EnvTransition) -> EnvTransition: - transition = transition.copy() - - state = transition.get(TransitionKey.OBSERVATION, {}).get(OBS_STATE) + def complementary_data(self, complementary_data: dict[str, Any]) -> dict[str, Any]: + state = (self.transition.get(TransitionKey.OBSERVATION) or {}).get(OBS_STATE) if state is None: raise ValueError("State is required for PI05") - tasks = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get(self.task_key) + tasks = complementary_data.get(self.task_key) if tasks is None: raise ValueError("No task found in complementary data") - # TODO: check if this necessary - state = deepcopy(state) - # State should already be normalized to [-1, 1] by the NormalizerProcessorStep that runs before this step # Discretize into 256 bins (see openpi `PaligemmaTokenizer.tokenize()`) state_np = state.cpu().numpy() @@ -74,10 +69,11 @@ class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep): full_prompt = f"Task: {cleaned_text}, State: {state_str};\nAction: " full_prompts.append(full_prompt) - transition[TransitionKey.COMPLEMENTARY_DATA][self.task_key] = full_prompts - # Normalize state to [-1, 1] range if needed (assuming it's already normalized by normalizer processor step!!) - # Discretize into 256 bins (see openpi `PaligemmaTokenizer.tokenize()`) - return transition + complementary_data[self.task_key] = full_prompts + return complementary_data + + def get_config(self) -> dict[str, Any]: + return {"task_key": self.task_key, "max_state_dim": self.max_state_dim} def transform_features( self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] diff --git a/src/lerobot/policies/pi0_fast/processor_pi0_fast.py b/src/lerobot/policies/pi0_fast/processor_pi0_fast.py index 864dafcfb..7b5be8914 100644 --- a/src/lerobot/policies/pi0_fast/processor_pi0_fast.py +++ b/src/lerobot/policies/pi0_fast/processor_pi0_fast.py @@ -14,7 +14,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -from copy import deepcopy from dataclasses import dataclass from typing import Any @@ -22,10 +21,11 @@ import numpy as np import torch from lerobot.configs import PipelineFeatureType, PolicyFeature -from lerobot.lerobot_types import EnvTransition, TransitionKey +from lerobot.lerobot_types import TransitionKey from lerobot.processor import ( AbsoluteActionsProcessorStep, ActionTokenizerProcessorStep, + ComplementaryDataProcessorStep, PolicyAction, PolicyProcessorPipeline, ProcessorStep, @@ -42,7 +42,7 @@ from .configuration_pi0_fast import PI0FastConfig @ProcessorStepRegistry.register(name="pi0_fast_prepare_state_tokenizer_processor_step") @dataclass -class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep): +class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ComplementaryDataProcessorStep): """ Processor step to prepare the state and tokenize the language input. """ @@ -50,19 +50,14 @@ class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep): max_state_dim: int = 32 task_key: str = "task" - def __call__(self, transition: EnvTransition) -> EnvTransition: - transition = transition.copy() - - state = transition.get(TransitionKey.OBSERVATION, {}).get(OBS_STATE) + def complementary_data(self, complementary_data: dict[str, Any]) -> dict[str, Any]: + state = (self.transition.get(TransitionKey.OBSERVATION) or {}).get(OBS_STATE) if state is None: raise ValueError("State is required for PI0Fast") - tasks = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get(self.task_key) + tasks = complementary_data.get(self.task_key) if tasks is None: raise ValueError("No task found in complementary data") - # TODO: check if this necessary - state = deepcopy(state) - # State should already be normalized to [-1, 1] by the NormalizerProcessorStep that runs before this step # Discretize into 256 bins (see openpi `PaligemmaTokenizer.tokenize()`) state_np = state.cpu().numpy() @@ -75,10 +70,11 @@ class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep): full_prompt = f"Task: {cleaned_text}, State: {state_str};\n" full_prompts.append(full_prompt) - transition[TransitionKey.COMPLEMENTARY_DATA][self.task_key] = full_prompts - # Normalize state to [-1, 1] range if needed (assuming it's already normalized by normalizer processor step!!) - # Discretize into 256 bins (see openpi `PaligemmaTokenizer.tokenize()`) - return transition + complementary_data[self.task_key] = full_prompts + return complementary_data + + def get_config(self) -> dict[str, Any]: + return {"task_key": self.task_key, "max_state_dim": self.max_state_dim} def transform_features( self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] diff --git a/src/lerobot/policies/vla_jepa/processor_vla_jepa.py b/src/lerobot/policies/vla_jepa/processor_vla_jepa.py index aed2b754b..7d7457f2c 100644 --- a/src/lerobot/policies/vla_jepa/processor_vla_jepa.py +++ b/src/lerobot/policies/vla_jepa/processor_vla_jepa.py @@ -20,12 +20,11 @@ import torch from lerobot.policies.vla_jepa.configuration_vla_jepa import VLAJEPAConfig from lerobot.processor import ( - EnvTransition, PolicyAction, + PolicyActionProcessorStep, PolicyProcessorPipeline, ProcessorStep, ProcessorStepRegistry, - TransitionKey, UnnormalizerProcessorStep, make_default_policy_processor_steps, make_policy_processor_pipelines, @@ -33,22 +32,20 @@ from lerobot.processor import ( @ProcessorStepRegistry.register(name="vla_jepa_clip_actions") -class ClipActionsProcessorStep(ProcessorStep): +class ClipActionsProcessorStep(PolicyActionProcessorStep): """Clips action tensor to [-1, 1] before unnormalization.""" - def __call__(self, transition: EnvTransition) -> EnvTransition: - action = transition.get(TransitionKey.ACTION) - if action is not None: - transition = dict(transition) - transition[TransitionKey.ACTION] = action.clamp(-1.0, 1.0) - return transition + skip_if_missing = True + + def action(self, action: PolicyAction) -> PolicyAction: + return action.clamp(-1.0, 1.0) def transform_features(self, features): return features @ProcessorStepRegistry.register(name="vla_jepa_pre_snap_gripper") -class PreSnapGripperProcessorStep(ProcessorStep): +class PreSnapGripperProcessorStep(PolicyActionProcessorStep): """Snaps a gripper dimension to {0, 1} BEFORE unnormalization. Mirrors the original starVLA LIBERO eval: @@ -58,43 +55,49 @@ class PreSnapGripperProcessorStep(ProcessorStep): space where 0=open and 1=close. """ + skip_if_missing = True + def __init__(self, gripper_dim: int = 6, threshold: float = 0.5): self.gripper_dim = gripper_dim self.threshold = threshold - def __call__(self, transition: EnvTransition) -> EnvTransition: - action = transition.get(TransitionKey.ACTION) - if action is not None and action.shape[-1] > self.gripper_dim: - transition = dict(transition) - a = action.clone() - a[..., self.gripper_dim] = (a[..., self.gripper_dim] >= self.threshold).float() - transition[TransitionKey.ACTION] = a - return transition + def action(self, action: PolicyAction) -> PolicyAction: + if action.shape[-1] <= self.gripper_dim: + return action + a = action.clone() + a[..., self.gripper_dim] = (a[..., self.gripper_dim] >= self.threshold).float() + return a + + def get_config(self) -> dict[str, Any]: + return {"gripper_dim": self.gripper_dim, "threshold": self.threshold} def transform_features(self, features): return features @ProcessorStepRegistry.register(name="vla_jepa_binarize_gripper") -class BinarizeGripperProcessorStep(ProcessorStep): +class BinarizeGripperProcessorStep(PolicyActionProcessorStep): """Binarizes a gripper dimension after unnormalization. Maps continuous value to {-1, 1}: > threshold → -1, <= threshold → 1 (matches starVLA convention). Only applied when action has more dimensions than gripper_dim. """ + skip_if_missing = True + def __init__(self, gripper_dim: int = 6, threshold: float = 0.5): self.gripper_dim = gripper_dim self.threshold = threshold - def __call__(self, transition: EnvTransition) -> EnvTransition: - action = transition.get(TransitionKey.ACTION) - if action is not None and action.shape[-1] > self.gripper_dim: - transition = dict(transition) - a = action.clone() - a[..., self.gripper_dim] = 1.0 - 2.0 * (a[..., self.gripper_dim] > self.threshold).float() - transition[TransitionKey.ACTION] = a - return transition + def action(self, action: PolicyAction) -> PolicyAction: + if action.shape[-1] <= self.gripper_dim: + return action + a = action.clone() + a[..., self.gripper_dim] = 1.0 - 2.0 * (a[..., self.gripper_dim] > self.threshold).float() + return a + + def get_config(self) -> dict[str, Any]: + return {"gripper_dim": self.gripper_dim, "threshold": self.threshold} def transform_features(self, features): return features diff --git a/src/lerobot/policies/xvla/processor_xvla.py b/src/lerobot/policies/xvla/processor_xvla.py index 4ee8fe793..683a7c876 100644 --- a/src/lerobot/policies/xvla/processor_xvla.py +++ b/src/lerobot/policies/xvla/processor_xvla.py @@ -21,10 +21,12 @@ import numpy as np import torch from lerobot.configs import PipelineFeatureType, PolicyFeature -from lerobot.lerobot_types import EnvTransition, TransitionKey +from lerobot.lerobot_types import TransitionKey from lerobot.processor import ( + ComplementaryDataProcessorStep, ObservationProcessorStep, PolicyAction, + PolicyActionProcessorStep, PolicyProcessorPipeline, ProcessorStep, ProcessorStepRegistry, @@ -201,7 +203,7 @@ class LiberoProcessorStep(ObservationProcessorStep): @dataclass @ProcessorStepRegistry.register(name="xvla_image_scale") -class XVLAImageScaleProcessorStep(ProcessorStep): +class XVLAImageScaleProcessorStep(ObservationProcessorStep): """Scale image observations by 255 to convert from [0, 1] to [0, 255] range. This processor step multiplies all image observations by 255, which is required @@ -214,29 +216,22 @@ class XVLAImageScaleProcessorStep(ProcessorStep): image_keys: list[str] | None = None - def __call__(self, transition: EnvTransition) -> EnvTransition: + skip_if_missing = True + + def observation(self, observation): """Scale image observations by 255.""" - new_transition = transition.copy() - obs = new_transition.get(TransitionKey.OBSERVATION, {}) - if obs is None: - return new_transition - - # Make a copy of observations to avoid modifying the original - obs = obs.copy() - # Determine which keys to scale keys_to_scale = self.image_keys if keys_to_scale is None: # Auto-detect image keys - keys_to_scale = [k for k in obs if k.startswith(OBS_IMAGES)] + keys_to_scale = [k for k in observation if k.startswith(OBS_IMAGES)] # Scale each image for key in keys_to_scale: - if key in obs and isinstance(obs[key], torch.Tensor): - obs[key] = obs[key] * 255 + if key in observation and isinstance(observation[key], torch.Tensor): + observation[key] = observation[key] * 255 - new_transition[TransitionKey.OBSERVATION] = obs - return new_transition + return observation def transform_features(self, features): """Image scaling doesn't change feature structure.""" @@ -251,7 +246,7 @@ class XVLAImageScaleProcessorStep(ProcessorStep): @dataclass @ProcessorStepRegistry.register(name="xvla_image_to_float") -class XVLAImageToFloatProcessorStep(ProcessorStep): +class XVLAImageToFloatProcessorStep(ObservationProcessorStep): """Convert image observations from [0, 255] to [0, 1] range. This processor step divides image observations by 255 to convert from uint8-like @@ -270,32 +265,26 @@ class XVLAImageToFloatProcessorStep(ProcessorStep): image_keys: list[str] | None = None validate_range: bool = True - def __call__(self, transition: EnvTransition) -> EnvTransition: + skip_if_missing = True + + def observation(self, observation): """Convert image observations from [0, 255] to [0, 1].""" - new_transition = transition.copy() - obs = new_transition.get(TransitionKey.OBSERVATION, {}) - if obs is None: - return new_transition - - # Make a copy of observations to avoid modifying the original - obs = obs.copy() - # Determine which keys to convert keys_to_convert = self.image_keys if keys_to_convert is None: # Auto-detect image keys - keys_to_convert = [k for k in obs if k.startswith(OBS_IMAGES)] + keys_to_convert = [k for k in observation if k.startswith(OBS_IMAGES)] # Convert each image for key in keys_to_convert: - if key in obs and isinstance(obs[key], torch.Tensor): - tensor = obs[key] + if key in observation and isinstance(observation[key], torch.Tensor): + tensor = observation[key] min_val = tensor.min().item() max_val = tensor.max().item() if max_val <= 1.0: - obs[key] = tensor.float() # ensure float dtype, but no division + observation[key] = tensor.float() # ensure float dtype, but no division continue # Validate that values are in [0, 255] range if requested if self.validate_range and (min_val < 0.0 or max_val > 255.0): @@ -306,10 +295,9 @@ class XVLAImageToFloatProcessorStep(ProcessorStep): ) # Convert to float and divide by 255 - obs[key] = tensor.float() / 255.0 + observation[key] = tensor.float() / 255.0 - new_transition[TransitionKey.OBSERVATION] = obs - return new_transition + return observation def transform_features(self, features): """Image conversion doesn't change feature structure.""" @@ -325,7 +313,7 @@ class XVLAImageToFloatProcessorStep(ProcessorStep): @dataclass @ProcessorStepRegistry.register(name="xvla_imagenet_normalize") -class XVLAImageNetNormalizeProcessorStep(ProcessorStep): +class XVLAImageNetNormalizeProcessorStep(ObservationProcessorStep): """Normalize image observations using ImageNet statistics. This processor step applies ImageNet normalization (mean and std) to image observations. @@ -343,26 +331,20 @@ class XVLAImageNetNormalizeProcessorStep(ProcessorStep): image_keys: list[str] | None = None - def __call__(self, transition: EnvTransition) -> EnvTransition: + skip_if_missing = True + + def observation(self, observation): """Normalize image observations using ImageNet statistics.""" - new_transition = transition.copy() - obs = new_transition.get(TransitionKey.OBSERVATION, {}) - if obs is None: - return new_transition - - # Make a copy of observations to avoid modifying the original - obs = obs.copy() - # Determine which keys to normalize keys_to_normalize = self.image_keys if keys_to_normalize is None: # Auto-detect image keys - keys_to_normalize = [k for k in obs if k.startswith(OBS_IMAGES)] + keys_to_normalize = [k for k in observation if k.startswith(OBS_IMAGES)] # Normalize each image for key in keys_to_normalize: - if key in obs and isinstance(obs[key], torch.Tensor): - tensor = obs[key] + if key in observation and isinstance(observation[key], torch.Tensor): + tensor = observation[key] # Validate that values are in [0, 1] range min_val = tensor.min().item() @@ -384,10 +366,9 @@ class XVLAImageNetNormalizeProcessorStep(ProcessorStep): std = std.unsqueeze(0) # Normalize: (image - mean) / std - obs[key] = (tensor - mean) / std + observation[key] = (tensor - mean) / std - new_transition[TransitionKey.OBSERVATION] = obs - return new_transition + return observation def transform_features(self, features): """ImageNet normalization doesn't change feature structure.""" @@ -402,38 +383,32 @@ class XVLAImageNetNormalizeProcessorStep(ProcessorStep): @dataclass @ProcessorStepRegistry.register(name="xvla_add_domain_id") -class XVLAAddDomainIdProcessorStep(ProcessorStep): +class XVLAAddDomainIdProcessorStep(ComplementaryDataProcessorStep): """Add domain_id to complementary data. This processor step adds a domain_id tensor to the complementary data, which is used by XVLA to identify different robot embodiments or task domains. Args: - domain_id: The domain ID to add (default: 3) + domain_id: The domain ID to add (default: 0) """ domain_id: int = 0 - def __call__(self, transition: EnvTransition) -> EnvTransition: + def complementary_data(self, complementary_data): """Add domain_id to complementary data.""" - new_transition = transition.copy() - comp = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) - comp = {} if comp is None else comp.copy() - # Infer batch size from observation tensors - obs = new_transition.get(TransitionKey.OBSERVATION, {}) + obs = self.transition.get(TransitionKey.OBSERVATION) or {} batch_size = 1 - if obs: - for v in obs.values(): - if isinstance(v, torch.Tensor): - batch_size = v.shape[0] - break + for v in obs.values(): + if isinstance(v, torch.Tensor): + batch_size = v.shape[0] + break # Add domain_id tensor - comp["domain_id"] = torch.tensor([int(self.domain_id)] * batch_size, dtype=torch.long) + complementary_data["domain_id"] = torch.tensor([int(self.domain_id)] * batch_size, dtype=torch.long) - new_transition[TransitionKey.COMPLEMENTARY_DATA] = comp - return new_transition + return complementary_data def transform_features(self, features): """Domain ID addition doesn't change feature structure.""" @@ -448,7 +423,7 @@ class XVLAAddDomainIdProcessorStep(ProcessorStep): @dataclass @ProcessorStepRegistry.register(name="xvla_rotation_6d_to_axis_angle") -class XVLARotation6DToAxisAngleProcessorStep(ProcessorStep): +class XVLARotation6DToAxisAngleProcessorStep(PolicyActionProcessorStep): """Convert 6D rotation representation to axis-angle and reorganize action dimensions. This processor step takes actions with 6D rotation representation and converts them to @@ -465,14 +440,10 @@ class XVLARotation6DToAxisAngleProcessorStep(ProcessorStep): expected_action_dim: int = 10 - def __call__(self, transition: EnvTransition) -> EnvTransition: + skip_if_missing = True + + def action(self, action: PolicyAction) -> PolicyAction: """Convert 6D rotation to axis-angle in action.""" - new_transition = transition.copy() - action = new_transition.get(TransitionKey.ACTION) - - if action is None or not isinstance(action, torch.Tensor): - return new_transition - # Convert to numpy for processing device = action.device dtype = action.dtype @@ -494,10 +465,7 @@ class XVLARotation6DToAxisAngleProcessorStep(ProcessorStep): action_np[:, -1] = np.where(action_np[:, -1] > 0.5, 1.0, -1.0) # Convert back to tensor - action = torch.from_numpy(action_np).to(device=device, dtype=dtype) - - new_transition[TransitionKey.ACTION] = action - return new_transition + return torch.from_numpy(action_np).to(device=device, dtype=dtype) def transform_features(self, features): """Rotation conversion changes action dimension from 10 to 7.""" diff --git a/src/lerobot/processor/__init__.py b/src/lerobot/processor/__init__.py index 6ed762b6a..e780578ba 100644 --- a/src/lerobot/processor/__init__.py +++ b/src/lerobot/processor/__init__.py @@ -53,13 +53,15 @@ from .factory import ( ) from .gym_action_processor import ( Numpy2TorchActionProcessorStep, + Numpy2TorchTeleopActionProcessorStep, Torch2NumpyActionProcessorStep, ) from .hil_processor import ( AddTeleopActionAsComplimentaryDataStep, AddTeleopEventsAsInfoStep, GripperPenaltyProcessorStep, - GymHILAdapterProcessorStep, + GymHILInfoAdapterStep, + GymHILTeleopDataAdapterStep, ImageCropResizeProcessorStep, InterventionActionProcessorStep, RewardClassifierProcessorStep, @@ -126,7 +128,8 @@ __all__ = [ "DoneProcessorStep", "EnvAction", "EnvTransition", - "GymHILAdapterProcessorStep", + "GymHILInfoAdapterStep", + "GymHILTeleopDataAdapterStep", "GripperPenaltyProcessorStep", "hotswap_stats", "IdentityProcessorStep", @@ -148,6 +151,7 @@ __all__ = [ "NewLineTaskProcessorStep", "NormalizerProcessorStep", "Numpy2TorchActionProcessorStep", + "Numpy2TorchTeleopActionProcessorStep", "ObservationProcessorStep", "PolicyAction", "PolicyActionProcessorStep", diff --git a/src/lerobot/processor/gym_action_processor.py b/src/lerobot/processor/gym_action_processor.py index 549e2b1e0..a57d68350 100644 --- a/src/lerobot/processor/gym_action_processor.py +++ b/src/lerobot/processor/gym_action_processor.py @@ -17,11 +17,11 @@ from dataclasses import dataclass from lerobot.configs import PipelineFeatureType, PolicyFeature -from lerobot.lerobot_types import EnvAction, EnvTransition, PolicyAction, TransitionKey +from lerobot.lerobot_types import EnvAction, PolicyAction from .converters import to_tensor from .hil_processor import TELEOP_ACTION_KEY -from .pipeline import ActionProcessorStep, ProcessorStep, ProcessorStepRegistry +from .pipeline import ActionProcessorStep, ComplementaryDataProcessorStep, ProcessorStepRegistry @ProcessorStepRegistry.register("torch2numpy_action_processor") @@ -70,32 +70,36 @@ class Torch2NumpyActionProcessorStep(ActionProcessorStep): @ProcessorStepRegistry.register("numpy2torch_action_processor") @dataclass -class Numpy2TorchActionProcessorStep(ProcessorStep): +class Numpy2TorchActionProcessorStep(ActionProcessorStep): """Converts a NumPy array action to a PyTorch tensor when action is present.""" - def __call__(self, transition: EnvTransition) -> EnvTransition: - """Converts numpy action to torch tensor if action exists, otherwise passes through.""" - self._current_transition = transition.copy() - new_transition = self._current_transition + skip_if_missing = True - action = new_transition.get(TransitionKey.ACTION) - if action is not None: - if not isinstance(action, EnvAction): - raise TypeError( - f"Expected np.ndarray or None, got {type(action).__name__}. " - "Use appropriate processor for non-tensor actions." - ) - torch_action = to_tensor(action, dtype=None) # Preserve original dtype - new_transition[TransitionKey.ACTION] = torch_action - - complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) - if TELEOP_ACTION_KEY in complementary_data: - teleop_action = complementary_data[TELEOP_ACTION_KEY] - if isinstance(teleop_action, EnvAction): - complementary_data[TELEOP_ACTION_KEY] = to_tensor(teleop_action) - new_transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data - - return new_transition + def action(self, action: EnvAction) -> PolicyAction: + if not isinstance(action, EnvAction): + raise TypeError( + f"Expected np.ndarray or None, got {type(action).__name__}. " + "Use appropriate processor for non-tensor actions." + ) + return to_tensor(action, dtype=None) # Preserve original dtype + + def transform_features( + self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] + ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: + return features + + +@ProcessorStepRegistry.register("numpy2torch_teleop_action_processor") +@dataclass +class Numpy2TorchTeleopActionProcessorStep(ComplementaryDataProcessorStep): + """Converts a NumPy teleop action in the complementary data to a PyTorch tensor.""" + + def complementary_data(self, complementary_data: dict) -> dict: + if TELEOP_ACTION_KEY in complementary_data: + teleop_action = complementary_data[TELEOP_ACTION_KEY] + if isinstance(teleop_action, EnvAction): + complementary_data[TELEOP_ACTION_KEY] = to_tensor(teleop_action) + return complementary_data def transform_features( self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] diff --git a/src/lerobot/processor/hil_processor.py b/src/lerobot/processor/hil_processor.py index e0c9fe733..4470f7566 100644 --- a/src/lerobot/processor/hil_processor.py +++ b/src/lerobot/processor/hil_processor.py @@ -312,20 +312,39 @@ class TimeLimitProcessorStep(TruncatedProcessorStep): return features -@ProcessorStepRegistry.register("gym_hil_adapter_processor") -class GymHILAdapterProcessorStep(ProcessorStep): +@ProcessorStepRegistry.register("gym_hil_info_adapter") +class GymHILInfoAdapterStep(InfoProcessorStep): """ - Adapts the output of the `gym-hil` environment to the format expected by `lerobot` processors. + Adapts the `info` dictionary of the `gym-hil` environment to the format expected by + `lerobot` processors. - This step normalizes the `transition` object by: - 1. Copying `teleop_action` from `info` to `complementary_data`. - 2. Copying `is_intervention` from `info` (using the string key) to `info` (using the enum key). - 3. Copying `discrete_penalty` from `info` to `complementary_data`. + Mirrors `is_intervention` from the string key to the `TeleopEvents.IS_INTERVENTION` + enum key when present. """ - def __call__(self, transition: EnvTransition) -> EnvTransition: - info = transition.get(TransitionKey.INFO, {}) - complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) + def info(self, info: dict) -> dict: + if "is_intervention" in info: + info[TeleopEvents.IS_INTERVENTION] = info["is_intervention"] + return info + + def transform_features( + self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] + ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: + return features + + +@ProcessorStepRegistry.register("gym_hil_teleop_data_adapter") +class GymHILTeleopDataAdapterStep(ComplementaryDataProcessorStep): + """ + Copies teleoperation data emitted by the `gym-hil` environment from `info` into the + transition's complementary data. + + Copies `teleop_action` and `discrete_penalty` from `info` to `complementary_data` + when present. + """ + + def complementary_data(self, complementary_data: dict) -> dict: + info = self.transition.get(TransitionKey.INFO) or {} if TELEOP_ACTION_KEY in info: complementary_data[TELEOP_ACTION_KEY] = info[TELEOP_ACTION_KEY] @@ -333,13 +352,7 @@ class GymHILAdapterProcessorStep(ProcessorStep): if DISCRETE_PENALTY_KEY in info: complementary_data[DISCRETE_PENALTY_KEY] = info[DISCRETE_PENALTY_KEY] - if "is_intervention" in info: - info[TeleopEvents.IS_INTERVENTION] = info["is_intervention"] - - transition[TransitionKey.INFO] = info - transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data - - return transition + return complementary_data def transform_features( self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] @@ -349,7 +362,7 @@ class GymHILAdapterProcessorStep(ProcessorStep): @dataclass @ProcessorStepRegistry.register("gripper_penalty_processor") -class GripperPenaltyProcessorStep(ProcessorStep): +class GripperPenaltyProcessorStep(ComplementaryDataProcessorStep): """ Applies a small per-transition cost on the discrete gripper action. @@ -370,31 +383,30 @@ class GripperPenaltyProcessorStep(ProcessorStep): open_threshold: float = 0.1 closed_threshold: float = 0.9 - def __call__(self, transition: EnvTransition) -> EnvTransition: + def complementary_data(self, complementary_data: dict) -> dict: """ Calculates the gripper penalty and adds it to the complementary data. Args: - transition: The incoming environment transition. + complementary_data: The incoming complementary data dictionary. Returns: - The modified transition with the penalty added to complementary data. + The complementary data with the penalty added under the + `discrete_penalty` key. """ - new_transition = transition.copy() - action = new_transition.get(TransitionKey.ACTION) - complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) + action = self.transition.get(TransitionKey.ACTION) raw_joint_positions = complementary_data.get("raw_joint_positions") if raw_joint_positions is None: - return new_transition + return complementary_data current_gripper_pos = raw_joint_positions.get(f"{GRIPPER_KEY}.pos", None) if current_gripper_pos is None: - return new_transition + return complementary_data # During reset, the transition may not carry any action yet. if action is None: - return new_transition + return complementary_data # Gripper action is expected as the last action dimension. gripper_action = action[-1].item() @@ -414,12 +426,8 @@ class GripperPenaltyProcessorStep(ProcessorStep): gripper_penalty = self.penalty * int(gripper_penalty_bool) - # Update complementary data with penalty info - new_complementary_data = dict(complementary_data) - new_complementary_data[DISCRETE_PENALTY_KEY] = gripper_penalty - new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data - - return new_transition + complementary_data[DISCRETE_PENALTY_KEY] = gripper_penalty + return complementary_data def get_config(self) -> dict[str, Any]: """ @@ -436,10 +444,6 @@ class GripperPenaltyProcessorStep(ProcessorStep): "closed_threshold": self.closed_threshold, } - def reset(self) -> None: - """Resets the processor's internal state.""" - pass - def transform_features( self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: diff --git a/src/lerobot/processor/pipeline.py b/src/lerobot/processor/pipeline.py index 54036f986..3fd1b6b8e 100644 --- a/src/lerobot/processor/pipeline.py +++ b/src/lerobot/processor/pipeline.py @@ -38,7 +38,7 @@ from collections.abc import Callable, Iterable, Sequence from copy import deepcopy from dataclasses import dataclass, field from pathlib import Path -from typing import Any, TypedDict, TypeVar, cast +from typing import Any, ClassVar, TypedDict, TypeVar, cast import torch from huggingface_hub import hf_hub_download @@ -159,6 +159,14 @@ class ProcessorStep(ABC): _current_transition: EnvTransition | None = None + # Consulted by the specialized single-field bases (ObservationProcessorStep, ActionProcessorStep, + # etc.): when True, the step is skipped (the transition is returned unchanged) if its target field + # is None, instead of raising a ValueError. Set it as a plain class attribute in subclasses + # (`skip_if_missing = True`) so that dataclass steps don't pick it up as a field. Use it for steps + # that must tolerate partial transitions, e.g. action steps in a preprocessor that also runs at + # inference time (where the action is None) or steps in RL pipelines that run on reset transitions. + skip_if_missing: ClassVar[bool] = False + @property def transition(self) -> EnvTransition: """Provides access to the most recent transition being processed. @@ -1753,7 +1761,12 @@ PolicyProcessorPipeline = DataProcessorPipeline[TInput, TOutput] class ObservationProcessorStep(ProcessorStep, ABC): - """An abstract `ProcessorStep` that specifically targets the observation in a transition.""" + """An abstract `ProcessorStep` that specifically targets the observation in a transition. + + The `observation` hook may read other parts of the transition via `self.transition`, but only the + observation may be written. Set `skip_if_missing = True` on a subclass to skip the step (instead of + raising) when the transition carries no observation. + """ @abstractmethod def observation(self, observation: RobotObservation) -> RobotObservation: @@ -1773,6 +1786,8 @@ class ObservationProcessorStep(ProcessorStep, ABC): new_transition = self._current_transition observation = new_transition.get(TransitionKey.OBSERVATION) + if observation is None and self.skip_if_missing: + return new_transition if observation is None or not isinstance(observation, dict): raise ValueError("ObservationProcessorStep requires an observation in the transition.") @@ -1782,7 +1797,12 @@ class ObservationProcessorStep(ProcessorStep, ABC): class ActionProcessorStep(ProcessorStep, ABC): - """An abstract `ProcessorStep` that specifically targets the action in a transition.""" + """An abstract `ProcessorStep` that specifically targets the action in a transition. + + The `action` hook may read other parts of the transition via `self.transition`, but only the action + may be written. Set `skip_if_missing = True` on a subclass to skip the step (instead of raising) + when the transition carries no action, e.g. for steps in pipelines that also run at inference time. + """ @abstractmethod def action( @@ -1805,6 +1825,8 @@ class ActionProcessorStep(ProcessorStep, ABC): action = new_transition.get(TransitionKey.ACTION) if action is None: + if self.skip_if_missing: + return new_transition raise ValueError("ActionProcessorStep requires an action in the transition.") processed_action = self.action(action) @@ -1813,7 +1835,12 @@ class ActionProcessorStep(ProcessorStep, ABC): class RobotActionProcessorStep(ProcessorStep, ABC): - """An abstract `ProcessorStep` for processing a `RobotAction` (a dictionary).""" + """An abstract `ProcessorStep` for processing a `RobotAction` (a dictionary). + + The `action` hook may read other parts of the transition via `self.transition`, but only the action + may be written. Set `skip_if_missing = True` on a subclass to skip the step (instead of raising) + when the transition carries no action. + """ @abstractmethod def action(self, action: RobotAction) -> RobotAction: @@ -1833,6 +1860,8 @@ class RobotActionProcessorStep(ProcessorStep, ABC): new_transition = self._current_transition action = new_transition.get(TransitionKey.ACTION) + if action is None and self.skip_if_missing: + return new_transition if action is None or not isinstance(action, dict): raise ValueError(f"Action should be a RobotAction type (dict), but got {type(action)}") @@ -1842,7 +1871,12 @@ class RobotActionProcessorStep(ProcessorStep, ABC): class PolicyActionProcessorStep(ProcessorStep, ABC): - """An abstract `ProcessorStep` for processing a `PolicyAction` (a tensor or dict of tensors).""" + """An abstract `ProcessorStep` for processing a `PolicyAction` (a tensor). + + The `action` hook may read other parts of the transition via `self.transition`, but only the action + may be written. Set `skip_if_missing = True` on a subclass to skip the step (instead of raising) + when the transition carries no action, e.g. for steps in pipelines that also run at inference time. + """ @abstractmethod def action(self, action: PolicyAction) -> PolicyAction: @@ -1862,6 +1896,8 @@ class PolicyActionProcessorStep(ProcessorStep, ABC): new_transition = self._current_transition action = new_transition.get(TransitionKey.ACTION) + if action is None and self.skip_if_missing: + return new_transition if not isinstance(action, PolicyAction): raise ValueError(f"Action should be a PolicyAction type (tensor), but got {type(action)}") @@ -1871,7 +1907,11 @@ class PolicyActionProcessorStep(ProcessorStep, ABC): class RewardProcessorStep(ProcessorStep, ABC): - """An abstract `ProcessorStep` that specifically targets the reward in a transition.""" + """An abstract `ProcessorStep` that specifically targets the reward in a transition. + + Set `skip_if_missing = True` on a subclass to skip the step (instead of raising) when the + transition carries no reward. + """ @abstractmethod def reward(self, reward) -> float | torch.Tensor: @@ -1892,6 +1932,8 @@ class RewardProcessorStep(ProcessorStep, ABC): reward = new_transition.get(TransitionKey.REWARD) if reward is None: + if self.skip_if_missing: + return new_transition raise ValueError("RewardProcessorStep requires a reward in the transition.") processed_reward = self.reward(reward) @@ -1900,7 +1942,11 @@ class RewardProcessorStep(ProcessorStep, ABC): class DoneProcessorStep(ProcessorStep, ABC): - """An abstract `ProcessorStep` that specifically targets the 'done' flag in a transition.""" + """An abstract `ProcessorStep` that specifically targets the 'done' flag in a transition. + + Set `skip_if_missing = True` on a subclass to skip the step (instead of raising) when the + transition carries no 'done' flag. + """ @abstractmethod def done(self, done) -> bool | torch.Tensor: @@ -1921,6 +1967,8 @@ class DoneProcessorStep(ProcessorStep, ABC): done = new_transition.get(TransitionKey.DONE) if done is None: + if self.skip_if_missing: + return new_transition raise ValueError("DoneProcessorStep requires a done flag in the transition.") processed_done = self.done(done) @@ -1929,7 +1977,11 @@ class DoneProcessorStep(ProcessorStep, ABC): class TruncatedProcessorStep(ProcessorStep, ABC): - """An abstract `ProcessorStep` that specifically targets the 'truncated' flag in a transition.""" + """An abstract `ProcessorStep` that specifically targets the 'truncated' flag in a transition. + + Set `skip_if_missing = True` on a subclass to skip the step (instead of raising) when the + transition carries no 'truncated' flag. + """ @abstractmethod def truncated(self, truncated) -> bool | torch.Tensor: @@ -1950,6 +2002,8 @@ class TruncatedProcessorStep(ProcessorStep, ABC): truncated = new_transition.get(TransitionKey.TRUNCATED) if truncated is None: + if self.skip_if_missing: + return new_transition raise ValueError("TruncatedProcessorStep requires a truncated flag in the transition.") processed_truncated = self.truncated(truncated) @@ -1958,7 +2012,11 @@ class TruncatedProcessorStep(ProcessorStep, ABC): class InfoProcessorStep(ProcessorStep, ABC): - """An abstract `ProcessorStep` that specifically targets the 'info' dictionary in a transition.""" + """An abstract `ProcessorStep` that specifically targets the 'info' dictionary in a transition. + + The `info` hook may read other parts of the transition via `self.transition`, but only the info + dictionary may be written. + """ @abstractmethod def info(self, info) -> dict[str, Any]: @@ -1978,6 +2036,8 @@ class InfoProcessorStep(ProcessorStep, ABC): new_transition = self._current_transition info = new_transition.get(TransitionKey.INFO) + if info is None and self.skip_if_missing: + return new_transition if info is None or not isinstance(info, dict): raise ValueError("InfoProcessorStep requires an info dictionary in the transition.") @@ -1987,7 +2047,11 @@ class InfoProcessorStep(ProcessorStep, ABC): class ComplementaryDataProcessorStep(ProcessorStep, ABC): - """An abstract `ProcessorStep` that targets the 'complementary_data' in a transition.""" + """An abstract `ProcessorStep` that targets the 'complementary_data' in a transition. + + The `complementary_data` hook may read other parts of the transition via `self.transition` (e.g. an + action or observation the step derives data from), but only the complementary data may be written. + """ @abstractmethod def complementary_data(self, complementary_data) -> dict[str, Any]: @@ -2007,6 +2071,8 @@ class ComplementaryDataProcessorStep(ProcessorStep, ABC): new_transition = self._current_transition complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA) + if complementary_data is None and self.skip_if_missing: + return new_transition if complementary_data is None or not isinstance(complementary_data, dict): raise ValueError("ComplementaryDataProcessorStep requires complementary data in the transition.") diff --git a/src/lerobot/processor/relative_action_processor.py b/src/lerobot/processor/relative_action_processor.py index 4fe007f79..1b3d2d92c 100644 --- a/src/lerobot/processor/relative_action_processor.py +++ b/src/lerobot/processor/relative_action_processor.py @@ -20,11 +20,11 @@ import torch from torch import Tensor from lerobot.configs import PipelineFeatureType, PolicyFeature -from lerobot.lerobot_types import EnvTransition, TransitionKey +from lerobot.lerobot_types import EnvTransition, PolicyAction, TransitionKey from lerobot.utils.constants import OBS_STATE from .delta_action_processor import MapDeltaActionToRobotActionStep, MapTensorToDeltaActionDictStep -from .pipeline import ProcessorStep, ProcessorStepRegistry +from .pipeline import PolicyActionProcessorStep, ProcessorStep, ProcessorStepRegistry # Re-export for backward compatibility __all__ = [ @@ -161,7 +161,7 @@ class RelativeActionsProcessorStep(ProcessorStep): @ProcessorStepRegistry.register("absolute_actions_processor") @dataclass -class AbsoluteActionsProcessorStep(ProcessorStep): +class AbsoluteActionsProcessorStep(PolicyActionProcessorStep): """Converts relative actions back to absolute actions (action += state) for all dimensions. Mirrors OpenPI's AbsoluteActions transform. Applied during postprocessing so @@ -176,9 +176,11 @@ class AbsoluteActionsProcessorStep(ProcessorStep): enabled: bool = False relative_step: RelativeActionsProcessorStep | None = field(default=None, repr=False) - def __call__(self, transition: EnvTransition) -> EnvTransition: + skip_if_missing = True + + def action(self, action: PolicyAction) -> PolicyAction: if not self.enabled: - return transition + return action if self.relative_step is None: raise RuntimeError( @@ -193,14 +195,8 @@ class AbsoluteActionsProcessorStep(ProcessorStep): "but no state has been cached. Ensure the preprocessor runs before the postprocessor." ) - new_transition = transition.copy() - action = new_transition.get(TransitionKey.ACTION) - if action is None: - return new_transition - mask = self.relative_step._build_mask(action.shape[-1]) - new_transition[TransitionKey.ACTION] = to_absolute_actions(action, cached_state, mask) - return new_transition + return to_absolute_actions(action, cached_state, mask) def get_config(self) -> dict[str, Any]: return {"enabled": self.enabled} diff --git a/src/lerobot/processor/tokenizer_processor.py b/src/lerobot/processor/tokenizer_processor.py index 967159144..07de314dd 100644 --- a/src/lerobot/processor/tokenizer_processor.py +++ b/src/lerobot/processor/tokenizer_processor.py @@ -41,7 +41,7 @@ from lerobot.utils.constants import ( ) from lerobot.utils.import_utils import _transformers_available -from .pipeline import ActionProcessorStep, ObservationProcessorStep, ProcessorStepRegistry +from .pipeline import ComplementaryDataProcessorStep, ObservationProcessorStep, ProcessorStepRegistry # Conditional import for type checking and lazy loading if TYPE_CHECKING or _transformers_available: @@ -325,19 +325,19 @@ class TokenizerProcessorStep(ObservationProcessorStep): @dataclass @ProcessorStepRegistry.register(name="action_tokenizer_processor") -class ActionTokenizerProcessorStep(ActionProcessorStep): +class ActionTokenizerProcessorStep(ComplementaryDataProcessorStep): """ Processor step to tokenize action data using a fast action tokenizer. - This step takes action tensors from an `EnvTransition`, tokenizes them using + This step reads the action tensor from the `EnvTransition`, tokenizes it using a Hugging Face `transformers` AutoProcessor (such as the Physical Intelligence "fast" tokenizer), - and returns the tokenized action. + and stores the resulting token IDs and mask in the transition's complementary data. Requires the `transformers` library to be installed. Attributes: - tokenizer_name: The name of a pretrained processor from the Hugging Face Hub (e.g., "lerobot/fast-action-tokenizer"). - tokenizer: A pre-initialized processor/tokenizer object. If provided, `tokenizer_name` is ignored. + action_tokenizer_name: The name of a pretrained processor from the Hugging Face Hub (e.g., "lerobot/fast-action-tokenizer"). + action_tokenizer_input_object: A pre-initialized processor/tokenizer object. If provided, `action_tokenizer_name` is ignored. trust_remote_code: Whether to trust remote code when loading the tokenizer (required for some tokenizers). action_tokenizer: The internal tokenizer/processor instance, loaded during initialization. paligemma_tokenizer_name: The name of a pretrained PaliGemma tokenizer from the Hugging Face Hub (e.g., "google/paligemma-3b-pt-224"). @@ -392,37 +392,27 @@ class ActionTokenizerProcessorStep(ActionProcessorStep): add_bos_token=False, ) - def __call__(self, transition: EnvTransition) -> EnvTransition: + def complementary_data(self, complementary_data: dict[str, Any]) -> dict[str, Any]: """ - Applies action tokenization to the transition. - - This overrides the base class to handle both tokens and mask. + Tokenizes the transition's action and adds the tokens and mask to the complementary data. Args: - transition: The input transition with action data. + complementary_data: The input complementary data dictionary. Returns: - The processed transition with tokenized actions and mask in complementary data. + The complementary data with tokenized actions and mask added. """ - self._current_transition = transition.copy() - new_transition = self._current_transition - - action = new_transition.get(TransitionKey.ACTION) + action = self.transition.get(TransitionKey.ACTION) if action is None: # During inference, no action is available, skip tokenization - return new_transition + return complementary_data # Tokenize and get both tokens and mask tokens, mask = self._tokenize_action(action) - # Store mask in complementary data - complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) - if complementary_data is None: - complementary_data = {} complementary_data[ACTION_TOKEN_MASK] = mask complementary_data[ACTION_TOKENS] = tokens - new_transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data - return new_transition + return complementary_data def _act_tokens_to_paligemma_tokens(self, tokens: torch.Tensor) -> torch.Tensor: """ @@ -529,14 +519,6 @@ class ActionTokenizerProcessorStep(ActionProcessorStep): return tokens_batch, masks_batch - def action(self, action: torch.Tensor) -> torch.Tensor: - """ - This method is not used since we override __call__. - Required by ActionProcessorStep ABC. - """ - tokens, _ = self._tokenize_action(action) - return tokens - def get_config(self) -> dict[str, Any]: """ Returns the serializable configuration of the processor. @@ -550,6 +532,8 @@ class ActionTokenizerProcessorStep(ActionProcessorStep): config = { "trust_remote_code": self.trust_remote_code, "max_action_tokens": self.max_action_tokens, + "fast_skip_tokens": self.fast_skip_tokens, + "paligemma_tokenizer_name": self.paligemma_tokenizer_name, } # Only save tokenizer_name if it was used to create the tokenizer @@ -562,15 +546,15 @@ class ActionTokenizerProcessorStep(ActionProcessorStep): self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: """ - Updates feature definitions to reflect tokenized actions. + Returns the policy features unchanged. - 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,). + The tokenized actions and mask are stored in complementary data, which is not + tracked in the policy features dictionary. Args: features: The dictionary of existing policy features. Returns: - The updated dictionary of policy features. + The dictionary of policy features, unchanged. """ return features diff --git a/src/lerobot/rewards/robometer/processor_robometer.py b/src/lerobot/rewards/robometer/processor_robometer.py index 764833202..f512ae9c6 100644 --- a/src/lerobot/rewards/robometer/processor_robometer.py +++ b/src/lerobot/rewards/robometer/processor_robometer.py @@ -25,13 +25,13 @@ from PIL import Image from torch import Tensor from lerobot.configs import PipelineFeatureType, PolicyFeature -from lerobot.lerobot_types import EnvTransition, TransitionKey +from lerobot.lerobot_types import TransitionKey from lerobot.processor import ( AddBatchDimensionProcessorStep, DeviceProcessorStep, + ObservationProcessorStep, PolicyAction, PolicyProcessorPipeline, - ProcessorStep, ProcessorStepRegistry, policy_action_to_transition, ) @@ -105,7 +105,7 @@ def _expand_tasks(task: Any, *, batch_size: int, default: str | None) -> list[st @dataclass @ProcessorStepRegistry.register(name="robometer_encoder") -class RobometerEncoderProcessorStep(ProcessorStep): +class RobometerEncoderProcessorStep(ObservationProcessorStep): """Encode raw frames + task into Qwen-VL tensors for the Robometer model. Loads a :class:`~transformers.AutoProcessor` matching ``base_model_id`` and @@ -160,11 +160,8 @@ class RobometerEncoderProcessorStep(ProcessorStep): if token not in tokenizer.get_vocab(): tokenizer.add_special_tokens({"additional_special_tokens": [token]}) - def __call__(self, transition: EnvTransition) -> EnvTransition: - observation = transition.get(TransitionKey.OBSERVATION) - complementary = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {} - if not isinstance(observation, dict): - raise ValueError("RobometerEncoderProcessorStep requires an observation dict") + def observation(self, observation: dict[str, Any]) -> dict[str, Any]: + complementary = self.transition.get(TransitionKey.COMPLEMENTARY_DATA) or {} if self.image_key not in observation: raise KeyError(f"Robometer expected image key {self.image_key!r} in observation") @@ -190,13 +187,9 @@ class RobometerEncoderProcessorStep(ProcessorStep): ] encoded = self.encode_samples(samples) - new_observation = dict(observation) for key, value in encoded.items(): - new_observation[f"{ROBOMETER_FEATURE_PREFIX}{key}"] = value - - new_transition = transition.copy() - new_transition[TransitionKey.OBSERVATION] = new_observation - return new_transition + observation[f"{ROBOMETER_FEATURE_PREFIX}{key}"] = value + return observation def encode_samples(self, samples: list[tuple[np.ndarray, str]]) -> dict[str, Tensor]: """Run the Qwen-VL processor on a list of ``(frames, task)`` samples.""" diff --git a/src/lerobot/rewards/sarm/processor_sarm.py b/src/lerobot/rewards/sarm/processor_sarm.py index d0597ccc1..6beb4be56 100644 --- a/src/lerobot/rewards/sarm/processor_sarm.py +++ b/src/lerobot/rewards/sarm/processor_sarm.py @@ -48,13 +48,13 @@ else: Faker = None # type: ignore[assignment, misc] from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature -from lerobot.lerobot_types import EnvTransition, PolicyAction, TransitionKey +from lerobot.lerobot_types import PolicyAction, TransitionKey from lerobot.processor import ( AddBatchDimensionProcessorStep, DeviceProcessorStep, NormalizerProcessorStep, + ObservationProcessorStep, PolicyProcessorPipeline, - ProcessorStep, RenameObservationsProcessorStep, from_tensor_to_numpy, policy_action_to_transition, @@ -73,7 +73,7 @@ from .sarm_utils import ( logger = logging.getLogger(__name__) -class SARMEncodingProcessorStep(ProcessorStep): +class SARMEncodingProcessorStep(ObservationProcessorStep): """ProcessorStep that encodes images and text with CLIP and generates stage and progress labels for SARM.""" def __init__( @@ -257,9 +257,9 @@ class SARMEncodingProcessorStep(ProcessorStep): return annotations - def __call__(self, transition: EnvTransition) -> EnvTransition: + def observation(self, observation: dict[str, Any]) -> dict[str, Any]: """ - Encode images, text, and normalize states in the transition. + Encode images, text, and normalize states in the observation. Implements SARM training data preparation: - Applies language perturbation (20% probability) @@ -267,9 +267,7 @@ class SARMEncodingProcessorStep(ProcessorStep): - Generates stage+tau targets for all frames - Outputs lengths tensor for valid sequence masking """ - new_transition = transition.copy() if hasattr(transition, "copy") else dict(transition) - observation = new_transition.get(TransitionKey.OBSERVATION) - comp_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) + comp_data = self.transition.get(TransitionKey.COMPLEMENTARY_DATA) or {} frame_index = comp_data.get("index") episode_index = comp_data.get("episode_index") @@ -392,8 +390,7 @@ class SARMEncodingProcessorStep(ProcessorStep): ) observation["dense_targets"] = dense_targets - new_transition[TransitionKey.OBSERVATION] = observation - return new_transition + return observation def _compute_batch_targets( self, diff --git a/src/lerobot/rewards/topreward/processor_topreward.py b/src/lerobot/rewards/topreward/processor_topreward.py index 75c5fd02a..0936da543 100644 --- a/src/lerobot/rewards/topreward/processor_topreward.py +++ b/src/lerobot/rewards/topreward/processor_topreward.py @@ -23,13 +23,13 @@ import torch from torch import Tensor from lerobot.configs import PipelineFeatureType, PolicyFeature -from lerobot.lerobot_types import EnvTransition, TransitionKey +from lerobot.lerobot_types import TransitionKey from lerobot.processor import ( AddBatchDimensionProcessorStep, DeviceProcessorStep, + ObservationProcessorStep, PolicyAction, PolicyProcessorPipeline, - ProcessorStep, ProcessorStepRegistry, policy_action_to_transition, ) @@ -107,7 +107,7 @@ def _expand_tasks(task: Any, *, batch_size: int, default: str | None) -> list[st @dataclass @ProcessorStepRegistry.register(name="topreward_encoder") -class TOPRewardEncoderProcessorStep(ProcessorStep): +class TOPRewardEncoderProcessorStep(ObservationProcessorStep): """Encode raw frames + task into Qwen-VL tensors for the TOPReward model. Loads a :class:`~transformers.AutoProcessor` matching ``vlm_name`` and @@ -142,9 +142,8 @@ class TOPRewardEncoderProcessorStep(ProcessorStep): require_package("transformers", extra="topreward") self._processor = AutoProcessor.from_pretrained(self.vlm_name, trust_remote_code=True) - def __call__(self, transition: EnvTransition) -> EnvTransition: - observation = transition.get(TransitionKey.OBSERVATION) - complementary = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {} + def observation(self, observation: dict[str, Any]) -> dict[str, Any]: + complementary = self.transition.get(TransitionKey.COMPLEMENTARY_DATA) or {} if self.image_key not in observation: raise KeyError(f"TOPReward expected image key {self.image_key!r} in observation") @@ -161,13 +160,9 @@ class TOPRewardEncoderProcessorStep(ProcessorStep): encoded = self._encode_batch(videos, tasks, batch_size) - new_observation = dict(observation) for key, value in encoded.items(): - new_observation[f"{TOPREWARD_FEATURE_PREFIX}{key}"] = value - - new_transition = transition.copy() - new_transition[TransitionKey.OBSERVATION] = new_observation - return new_transition + observation[f"{TOPREWARD_FEATURE_PREFIX}{key}"] = value + return observation def _encode_batch(self, videos: Tensor, tasks: list[str], batch_size) -> dict[str, Any]: """Tokenise a batch of (frames, task) pairs into Qwen-VL tensors. diff --git a/src/lerobot/rl/gym_manipulator.py b/src/lerobot/rl/gym_manipulator.py index 03f7b4eea..eef7c50ce 100644 --- a/src/lerobot/rl/gym_manipulator.py +++ b/src/lerobot/rl/gym_manipulator.py @@ -36,12 +36,14 @@ from lerobot.processor import ( DeviceProcessorStep, EnvTransition, GripperPenaltyProcessorStep, - GymHILAdapterProcessorStep, + GymHILInfoAdapterStep, + GymHILTeleopDataAdapterStep, ImageCropResizeProcessorStep, InterventionActionProcessorStep, MapDeltaActionToRobotActionStep, MapTensorToDeltaActionDictStep, Numpy2TorchActionProcessorStep, + Numpy2TorchTeleopActionProcessorStep, RewardClassifierProcessorStep, RobotActionToPolicyActionProcessorStep, RobotObservation, @@ -59,11 +61,12 @@ from lerobot.robots import ( # noqa: F401 ) from lerobot.robots.robot import Robot from lerobot.robots.so_follower.robot_kinematic_processor import ( + AddIKSolutionStep, EEBoundsAndSafety, EEReferenceAndDelta, ForwardKinematicsJointsToEEObservation, GripperVelocityToJoint, - InverseKinematicsRLStep, + InverseKinematicsEEToJoints, ) from lerobot.teleoperators import ( gamepad, # noqa: F401 @@ -382,8 +385,10 @@ def make_processors( ] env_pipeline_steps = [ - GymHILAdapterProcessorStep(), + GymHILInfoAdapterStep(), + GymHILTeleopDataAdapterStep(), Numpy2TorchActionProcessorStep(), + Numpy2TorchTeleopActionProcessorStep(), VanillaObservationProcessorStep(), ] @@ -494,6 +499,9 @@ def make_processors( # Replace InverseKinematicsProcessor with new kinematic processors if cfg.processor.inverse_kinematics is not None and kinematics_solver is not None: + ik_step = InverseKinematicsEEToJoints( + kinematics=kinematics_solver, motor_names=motor_names, initial_guess_current_joints=False + ) # Add EE bounds and safety processor inverse_kinematics_steps = [ MapTensorToDeltaActionDictStep( @@ -515,9 +523,8 @@ def make_processors( speed_factor=1.0, discrete_gripper=True, ), - InverseKinematicsRLStep( - kinematics=kinematics_solver, motor_names=motor_names, initial_guess_current_joints=False - ), + ik_step, + AddIKSolutionStep(ik_step=ik_step), ] action_pipeline_steps.extend(inverse_kinematics_steps) action_pipeline_steps.append(RobotActionToPolicyActionProcessorStep(motor_names=motor_names)) diff --git a/src/lerobot/robots/so_follower/robot_kinematic_processor.py b/src/lerobot/robots/so_follower/robot_kinematic_processor.py index 98bb54f67..0a4530851 100644 --- a/src/lerobot/robots/so_follower/robot_kinematic_processor.py +++ b/src/lerobot/robots/so_follower/robot_kinematic_processor.py @@ -23,9 +23,8 @@ import numpy as np from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature from lerobot.model import RobotKinematics from lerobot.processor import ( - EnvTransition, + ComplementaryDataProcessorStep, ObservationProcessorStep, - ProcessorStep, ProcessorStepRegistry, RobotAction, RobotActionProcessorStep, @@ -370,6 +369,34 @@ class InverseKinematicsEEToJoints(RobotActionProcessorStep): self.q_curr = None +@ProcessorStepRegistry.register("add_ik_solution") +@dataclass +class AddIKSolutionStep(ComplementaryDataProcessorStep): + """ + Records the latest IK solution in the transition's complementary data. + + Placed immediately after an `InverseKinematicsEEToJoints` step, it exposes that step's joint-space + solution under ``complementary_data["IK_solution"]`` so that downstream consumers (e.g. + `EEReferenceAndDelta` with ``use_ik_solution=True``) can reuse it. + + Attributes: + ik_step: The IK step whose solution state is recorded. + """ + + ik_step: InverseKinematicsEEToJoints + + def complementary_data(self, complementary_data: dict[str, Any]) -> dict[str, Any]: + if self.ik_step.q_curr is None: + return complementary_data + complementary_data["IK_solution"] = self.ik_step.q_curr.copy() + return complementary_data + + def transform_features( + self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] + ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: + return features + + @ProcessorStepRegistry.register("gripper_velocity_to_joint") @dataclass class GripperVelocityToJoint(RobotActionProcessorStep): @@ -522,127 +549,3 @@ class ForwardKinematicsJointsToEEAction(RobotActionProcessorStep): type=FeatureType.ACTION, shape=(1,) ) return features - - -@ProcessorStepRegistry.register(name="forward_kinematics_joints_to_ee") -@dataclass -class ForwardKinematicsJointsToEE(ProcessorStep): - kinematics: RobotKinematics - motor_names: list[str] - - def __post_init__(self): - self.joints_to_ee_action_processor = ForwardKinematicsJointsToEEAction( - kinematics=self.kinematics, motor_names=self.motor_names - ) - self.joints_to_ee_observation_processor = ForwardKinematicsJointsToEEObservation( - kinematics=self.kinematics, motor_names=self.motor_names - ) - - def __call__(self, transition: EnvTransition) -> EnvTransition: - if transition.get(TransitionKey.ACTION) is not None: - transition = self.joints_to_ee_action_processor(transition) - if transition.get(TransitionKey.OBSERVATION) is not None: - transition = self.joints_to_ee_observation_processor(transition) - return transition - - def transform_features( - self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] - ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: - if features[PipelineFeatureType.ACTION] is not None: - features = self.joints_to_ee_action_processor.transform_features(features) - if features[PipelineFeatureType.OBSERVATION] is not None: - features = self.joints_to_ee_observation_processor.transform_features(features) - return features - - -@ProcessorStepRegistry.register("inverse_kinematics_rl_step") -@dataclass -class InverseKinematicsRLStep(ProcessorStep): - """ - Computes desired joint positions from a target end-effector pose using inverse kinematics (IK). - - This is modified from the InverseKinematicsEEToJoints step to be used in the RL pipeline. - """ - - kinematics: RobotKinematics - motor_names: list[str] - q_curr: np.ndarray | None = field(default=None, init=False, repr=False) - initial_guess_current_joints: bool = True - - def __call__(self, transition: EnvTransition) -> EnvTransition: - new_transition = dict(transition) - action = new_transition.get(TransitionKey.ACTION) - if action is None: - raise ValueError("Action is required for InverseKinematicsEEToJoints") - action = dict(action) - - x = action.pop("ee.x") - y = action.pop("ee.y") - z = action.pop("ee.z") - wx = action.pop("ee.wx") - wy = action.pop("ee.wy") - wz = action.pop("ee.wz") - gripper_pos = action.pop("ee.gripper_pos") - - if None in (x, y, z, wx, wy, wz, gripper_pos): - raise ValueError( - "Missing required end-effector pose components: ee.x, ee.y, ee.z, ee.wx, ee.wy, ee.wz, ee.gripper_pos must all be present in action" - ) - - raw_observation = new_transition.get(TransitionKey.OBSERVATION) - if raw_observation is None: - raise ValueError("Joints observation is require for computing robot kinematics") - - observation = raw_observation.copy() - - q_raw = np.array( - [float(v) for k, v in observation.items() if isinstance(k, str) and k.endswith(".pos")], - dtype=float, - ) - if q_raw is None: - raise ValueError("Joints observation is require for computing robot kinematics") - - if self.initial_guess_current_joints: # Use current joints as initial guess - self.q_curr = q_raw - else: # Use previous ik solution as initial guess - if self.q_curr is None: - self.q_curr = q_raw - - # Build desired 4x4 transform from pos + rotvec (twist) - t_des = np.eye(4, dtype=float) - t_des[:3, :3] = Rotation.from_rotvec([wx, wy, wz]).as_matrix() - t_des[:3, 3] = [x, y, z] - - # Compute inverse kinematics - q_target = self.kinematics.inverse_kinematics(self.q_curr, t_des) - self.q_curr = q_target - - # TODO: This is sentitive to order of motor_names = q_target mapping - for i, name in enumerate(self.motor_names): - if name != "gripper": - action[f"{name}.pos"] = float(q_target[i]) - else: - action["gripper.pos"] = float(gripper_pos) - - new_transition[TransitionKey.ACTION] = action - complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) - complementary_data["IK_solution"] = q_target - new_transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data - return new_transition - - def transform_features( - self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] - ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: - for feat in ["x", "y", "z", "wx", "wy", "wz", "gripper_pos"]: - features[PipelineFeatureType.ACTION].pop(f"ee.{feat}", None) - - for name in self.motor_names: - features[PipelineFeatureType.ACTION][f"{name}.pos"] = PolicyFeature( - type=FeatureType.ACTION, shape=(1,) - ) - - return features - - def reset(self): - """Resets the initial guess for the IK solver.""" - self.q_curr = None diff --git a/tests/policies/molmoact2/test_molmoact2.py b/tests/policies/molmoact2/test_molmoact2.py index 41c0f235f..c25d7e6ce 100644 --- a/tests/policies/molmoact2/test_molmoact2.py +++ b/tests/policies/molmoact2/test_molmoact2.py @@ -955,10 +955,10 @@ def test_joint_frame_transform_noop_when_none(): state = torch.tensor([[10.0, -90.0, -120.0]]) state_transition = {TransitionKey.OBSERVATION: {OBS_STATE: state}} - assert state_step(state_transition) is state_transition + assert torch.equal(state_step(state_transition)[TransitionKey.OBSERVATION][OBS_STATE], state) action_transition = {TransitionKey.ACTION: state} - assert action_step(action_transition) is action_transition + assert torch.equal(action_step(action_transition)[TransitionKey.ACTION], state) def test_action_padding_marks_only_real_dimensions(): diff --git a/tests/processor/test_pipeline.py b/tests/processor/test_pipeline.py index 0e9746a63..9691b2372 100644 --- a/tests/processor/test_pipeline.py +++ b/tests/processor/test_pipeline.py @@ -2436,3 +2436,148 @@ def test_initial_camera_not_overridden_by_step_image(): key = f"{OBS_IMAGES}.front" assert key in out assert out[key]["shape"] == (240, 320, 3) # from the step, not from initial + + +class _SkipIfMissingMarkerMixin: + """Provides a pass-through transform_features so test steps only need to implement their field hook.""" + + def transform_features(self, features): + return features + + +def test_skip_if_missing_observation_step(): + from lerobot.processor import ObservationProcessorStep + + class TolerantObsStep(_SkipIfMissingMarkerMixin, ObservationProcessorStep): + skip_if_missing = True + + def observation(self, observation): + observation["ran"] = True + return observation + + class StrictObsStep(_SkipIfMissingMarkerMixin, ObservationProcessorStep): + def observation(self, observation): + return observation + + transition = create_transition(action=torch.zeros(2)) + out = TolerantObsStep()(transition) + assert out[TransitionKey.OBSERVATION] is None + with pytest.raises(ValueError, match="requires an observation"): + StrictObsStep()(transition) + + # With the field present, the hook still runs. + out = TolerantObsStep()(create_transition(observation={"x": 1})) + assert out[TransitionKey.OBSERVATION]["ran"] is True + + +def test_skip_if_missing_action_steps(): + from lerobot.processor import ( + ActionProcessorStep, + PolicyActionProcessorStep, + RobotActionProcessorStep, + ) + + class TolerantActionStep(_SkipIfMissingMarkerMixin, ActionProcessorStep): + skip_if_missing = True + + def action(self, action): + return action + 1 + + class TolerantPolicyActionStep(_SkipIfMissingMarkerMixin, PolicyActionProcessorStep): + skip_if_missing = True + + def action(self, action): + return action + 1 + + class TolerantRobotActionStep(_SkipIfMissingMarkerMixin, RobotActionProcessorStep): + skip_if_missing = True + + def action(self, action): + return {**action, "ran": True} + + action_less = create_transition(observation={"x": 1}) + for step in (TolerantActionStep(), TolerantPolicyActionStep(), TolerantRobotActionStep()): + out = step(action_less) + assert out[TransitionKey.ACTION] is None + + # Present action still processed. + out = TolerantPolicyActionStep()(create_transition(action=torch.zeros(2))) + assert torch.equal(out[TransitionKey.ACTION], torch.ones(2)) + + # A present-but-wrong-type action still raises, even with skip_if_missing. + with pytest.raises(ValueError, match="PolicyAction"): + TolerantPolicyActionStep()(create_transition(action={"joint": 1.0})) + with pytest.raises(ValueError, match="RobotAction"): + TolerantRobotActionStep()(create_transition(action=torch.zeros(2))) + + +def test_skip_if_missing_scalar_and_dict_steps(): + from lerobot.processor import ( + ComplementaryDataProcessorStep, + DoneProcessorStep, + InfoProcessorStep, + RewardProcessorStep, + TruncatedProcessorStep, + ) + + class TolerantRewardStep(_SkipIfMissingMarkerMixin, RewardProcessorStep): + skip_if_missing = True + + def reward(self, reward): + return reward + 1.0 + + class TolerantDoneStep(_SkipIfMissingMarkerMixin, DoneProcessorStep): + skip_if_missing = True + + def done(self, done): + return done + + class TolerantTruncatedStep(_SkipIfMissingMarkerMixin, TruncatedProcessorStep): + skip_if_missing = True + + def truncated(self, truncated): + return truncated + + class TolerantInfoStep(_SkipIfMissingMarkerMixin, InfoProcessorStep): + skip_if_missing = True + + def info(self, info): + return {**info, "ran": True} + + class TolerantComplementaryStep(_SkipIfMissingMarkerMixin, ComplementaryDataProcessorStep): + skip_if_missing = True + + def complementary_data(self, complementary_data): + return {**complementary_data, "ran": True} + + # create_transition never yields None for these fields, so build the partial transition by hand. + empty_transition: EnvTransition = {} + for step, key in ( + (TolerantRewardStep(), TransitionKey.REWARD), + (TolerantDoneStep(), TransitionKey.DONE), + (TolerantTruncatedStep(), TransitionKey.TRUNCATED), + (TolerantInfoStep(), TransitionKey.INFO), + (TolerantComplementaryStep(), TransitionKey.COMPLEMENTARY_DATA), + ): + out = step(empty_transition) + assert out.get(key) is None + + # With the fields present (create_transition defaults), the hooks still run. + full = create_transition() + assert TolerantRewardStep()(full)[TransitionKey.REWARD] == 1.0 + assert TolerantDoneStep()(full)[TransitionKey.DONE] is False + assert TolerantTruncatedStep()(full)[TransitionKey.TRUNCATED] is False + assert TolerantInfoStep()(full)[TransitionKey.INFO]["ran"] is True + assert TolerantComplementaryStep()(full)[TransitionKey.COMPLEMENTARY_DATA]["ran"] is True + + +def test_skip_if_missing_default_is_strict(): + from lerobot.processor import ActionProcessorStep + + class DefaultActionStep(_SkipIfMissingMarkerMixin, ActionProcessorStep): + def action(self, action): + return action + + assert DefaultActionStep.skip_if_missing is False + with pytest.raises(ValueError, match="requires an action"): + DefaultActionStep()(create_transition(observation={"x": 1})) diff --git a/tests/robots/test_robot_kinematic_processor.py b/tests/robots/test_robot_kinematic_processor.py index 7a8183c21..ca151e1a5 100644 --- a/tests/robots/test_robot_kinematic_processor.py +++ b/tests/robots/test_robot_kinematic_processor.py @@ -24,7 +24,6 @@ from lerobot.robots.so_follower.robot_kinematic_processor import ( ForwardKinematicsJointsToEEObservation, GripperVelocityToJoint, InverseKinematicsEEToJoints, - InverseKinematicsRLStep, ) MOTOR_NAMES = ["shoulder_pan", "shoulder_lift", "elbow_flex", "wrist_flex", "wrist_roll", "gripper"] @@ -62,13 +61,11 @@ EE_ACTION = dict.fromkeys(EE_KEYS, 0.0) ), (InverseKinematicsEEToJoints(kinematics=None, motor_names=MOTOR_NAMES), dict(EE_ACTION)), (GripperVelocityToJoint(), {**EE_ACTION, "ee.gripper_vel": 0.0}), - (InverseKinematicsRLStep(kinematics=None, motor_names=MOTOR_NAMES), dict(EE_ACTION)), ], ids=[ "ee_reference_and_delta", "inverse_kinematics_ee_to_joints", "gripper_velocity_to_joint", - "inverse_kinematics_rl_step", ], ) def test_missing_observation_raises_value_error(step, action): diff --git a/tests/robots/test_so_follower_kinematic_processor.py b/tests/robots/test_so_follower_kinematic_processor.py new file mode 100644 index 000000000..af1a7a1d4 --- /dev/null +++ b/tests/robots/test_so_follower_kinematic_processor.py @@ -0,0 +1,136 @@ +#!/usr/bin/env python + +# Copyright 2025 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from unittest.mock import MagicMock + +import numpy as np + +from lerobot.processor import TransitionKey +from lerobot.processor.converters import create_transition +from lerobot.robots.so_follower.robot_kinematic_processor import ( + AddIKSolutionStep, + ForwardKinematicsJointsToEEAction, + ForwardKinematicsJointsToEEObservation, + InverseKinematicsEEToJoints, +) + +MOTOR_NAMES = ["shoulder_pan", "shoulder_lift", "elbow_flex", "wrist_flex", "wrist_roll", "gripper"] +EE_KEYS = {f"ee.{k}" for k in ["x", "y", "z", "wx", "wy", "wz", "gripper_pos"]} + + +def _joints(gripper_pos: float = 7.0) -> dict[str, float]: + joints = {f"{n}.pos": float(i) for i, n in enumerate(MOTOR_NAMES) if n != "gripper"} + joints["gripper.pos"] = gripper_pos + return joints + + +def _fk_kinematics(translation: tuple[float, float, float]) -> MagicMock: + transform = np.eye(4, dtype=float) + transform[:3, 3] = translation + kinematics = MagicMock() + kinematics.forward_kinematics.return_value = transform + return kinematics + + +def test_forward_kinematics_observation_step(): + kinematics = _fk_kinematics(translation=(0.1, 0.2, 0.3)) + step = ForwardKinematicsJointsToEEObservation(kinematics=kinematics, motor_names=MOTOR_NAMES) + + transition = create_transition(observation=_joints(gripper_pos=42.0)) + result = step(transition) + + observation = result[TransitionKey.OBSERVATION] + assert set(observation) == EE_KEYS + assert observation["ee.x"] == 0.1 + assert observation["ee.y"] == 0.2 + assert observation["ee.z"] == 0.3 + assert observation["ee.wx"] == observation["ee.wy"] == observation["ee.wz"] == 0.0 + assert observation["ee.gripper_pos"] == 42.0 + assert result[TransitionKey.ACTION] is None + + (fk_input,) = kinematics.forward_kinematics.call_args.args + np.testing.assert_allclose(fk_input, [0.0, 1.0, 2.0, 3.0, 4.0, 42.0]) + + +def test_forward_kinematics_action_step(): + kinematics = _fk_kinematics(translation=(-0.5, 0.0, 0.25)) + step = ForwardKinematicsJointsToEEAction(kinematics=kinematics, motor_names=MOTOR_NAMES) + + transition = create_transition(action=_joints(gripper_pos=13.0)) + result = step(transition) + + action = result[TransitionKey.ACTION] + assert set(action) == EE_KEYS + assert action["ee.x"] == -0.5 + assert action["ee.y"] == 0.0 + assert action["ee.z"] == 0.25 + assert action["ee.gripper_pos"] == 13.0 + assert result[TransitionKey.OBSERVATION] is None + + +def test_inverse_kinematics_then_add_ik_solution(): + q_target = np.array([10.0, 20.0, 30.0, 40.0, 50.0, 60.0]) + kinematics = MagicMock() + kinematics.inverse_kinematics.return_value = q_target + + ik_step = InverseKinematicsEEToJoints(kinematics=kinematics, motor_names=MOTOR_NAMES) + add_step = AddIKSolutionStep(ik_step=ik_step) + + ee_action = { + "ee.x": 0.1, + "ee.y": 0.2, + "ee.z": 0.3, + "ee.wx": 0.0, + "ee.wy": 0.0, + "ee.wz": 0.0, + "ee.gripper_pos": 55.0, + } + transition = create_transition(action=ee_action, observation=_joints()) + result = add_step(ik_step(transition)) + + action = result[TransitionKey.ACTION] + for i, name in enumerate(MOTOR_NAMES): + expected = 55.0 if name == "gripper" else q_target[i] + assert action[f"{name}.pos"] == expected + assert not EE_KEYS & set(action) + + ik_solution = result[TransitionKey.COMPLEMENTARY_DATA]["IK_solution"] + np.testing.assert_allclose(ik_solution, q_target) + assert ik_solution is not ik_step.q_curr + + +def test_add_ik_solution_without_solution_leaves_data_unchanged(): + ik_step = InverseKinematicsEEToJoints(kinematics=None, motor_names=MOTOR_NAMES) + add_step = AddIKSolutionStep(ik_step=ik_step) + + transition = create_transition(complementary_data={"foo": "bar"}) + result = add_step(transition) + + assert result[TransitionKey.COMPLEMENTARY_DATA] == {"foo": "bar"} + + +def test_reset_clears_ik_solution_state(): + q_target = np.array([1.0, 2.0, 3.0, 4.0, 5.0, 6.0]) + kinematics = MagicMock() + kinematics.inverse_kinematics.return_value = q_target + ik_step = InverseKinematicsEEToJoints(kinematics=kinematics, motor_names=MOTOR_NAMES) + + ee_action = dict.fromkeys(EE_KEYS, 0.0) + ik_step(create_transition(action=ee_action, observation=_joints())) + assert ik_step.q_curr is not None + + ik_step.reset() + assert ik_step.q_curr is None