diff --git a/src/lerobot/rollout/context.py b/src/lerobot/rollout/context.py index aa69c8db5..265f69efe 100644 --- a/src/lerobot/rollout/context.py +++ b/src/lerobot/rollout/context.py @@ -46,6 +46,7 @@ from lerobot.processor import ( from lerobot.processor.relative_action_processor import RelativeActionsProcessorStep from lerobot.robots import make_robot_from_config from lerobot.teleoperators import Teleoperator, make_teleoperator_from_config +from lerobot.utils.constants import OBS_STATE from lerobot.utils.feature_utils import combine_feature_dicts, hw_to_dataset_features from .configs import BaseStrategyConfig, DAggerStrategyConfig, RolloutConfig @@ -109,6 +110,26 @@ def _resolve_action_key_order( return policy_action_names +def _align_relative_state_feature_order( + hw_features: dict[str, dict], policy_action_names: list[str] | None +) -> dict[str, dict]: + """Align policy-facing state with named relative-action dimensions.""" + if not policy_action_names or OBS_STATE not in hw_features: + return hw_features + + state_feature = hw_features[OBS_STATE] + state_names = state_feature.get("names") + if not state_names or len(state_names) != len(policy_action_names): + return hw_features + if set(state_names) != set(policy_action_names) or state_names == policy_action_names: + return hw_features + + aligned = dict(hw_features) + aligned[OBS_STATE] = {**state_feature, "names": list(policy_action_names)} + logger.info("Aligned relative-action state order with checkpoint action names") + return aligned + + # --------------------------------------------------------------------------- # Sub-contexts # --------------------------------------------------------------------------- @@ -431,10 +452,22 @@ def build_rollout_context( }, ) - if isinstance(cfg.inference, SyncInferenceConfig) and any( - isinstance(step, RelativeActionsProcessorStep) and step.enabled - for step in getattr(preprocessor, "steps", ()) - ): + relative_action_step = next( + ( + step + for step in getattr(preprocessor, "steps", ()) + if isinstance(step, RelativeActionsProcessorStep) and step.enabled + ), + None, + ) + if relative_action_step is not None: + relative_action_names = relative_action_step.action_names or policy_action_names + hw_features = _align_relative_state_feature_order( + hw_features, + list(relative_action_names) if relative_action_names else None, + ) + + if isinstance(cfg.inference, SyncInferenceConfig) and relative_action_step is not None: raise NotImplementedError( "SyncInferenceEngine does not support policies with relative actions for now." "Use --inference.type=rtc or remove relative action processor steps from the policy pipeline." diff --git a/tests/test_rollout.py b/tests/test_rollout.py index e0c5125b9..be5279e70 100644 --- a/tests/test_rollout.py +++ b/tests/test_rollout.py @@ -188,6 +188,42 @@ def test_trained_rtc_rollout_requires_capacity_for_max_delay(execution_horizon, _validate_trained_rtc_rollout_config(policy_config, inference_config) +def test_relative_state_order_follows_checkpoint_action_names(): + from lerobot.rollout.context import _align_relative_state_feature_order + from lerobot.utils.constants import OBS_STATE + from lerobot.utils.feature_utils import build_dataset_frame + + hw_features = { + OBS_STATE: { + "dtype": "float32", + "shape": (4,), + "names": ["left_joint.pos", "left_gripper.pos", "right_joint.pos", "right_gripper.pos"], + } + } + checkpoint_order = [ + "right_joint.pos", + "right_gripper.pos", + "left_joint.pos", + "left_gripper.pos", + ] + + aligned = _align_relative_state_feature_order(hw_features, checkpoint_order) + frame = build_dataset_frame( + aligned, + { + "left_joint.pos": 1.0, + "left_gripper.pos": 2.0, + "right_joint.pos": 3.0, + "right_gripper.pos": 4.0, + }, + prefix="observation", + ) + + assert aligned[OBS_STATE]["names"] == checkpoint_order + assert frame[OBS_STATE].tolist() == [3.0, 4.0, 1.0, 2.0] + assert hw_features[OBS_STATE]["names"][0] == "left_joint.pos" + + def test_sentry_config_defaults(): from lerobot.rollout import SentryStrategyConfig