mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-26 03:06:01 +00:00
fix(rollout): align relative state feature order
This commit is contained in:
@@ -46,6 +46,7 @@ from lerobot.processor import (
|
|||||||
from lerobot.processor.relative_action_processor import RelativeActionsProcessorStep
|
from lerobot.processor.relative_action_processor import RelativeActionsProcessorStep
|
||||||
from lerobot.robots import make_robot_from_config
|
from lerobot.robots import make_robot_from_config
|
||||||
from lerobot.teleoperators import Teleoperator, make_teleoperator_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 lerobot.utils.feature_utils import combine_feature_dicts, hw_to_dataset_features
|
||||||
|
|
||||||
from .configs import BaseStrategyConfig, DAggerStrategyConfig, RolloutConfig
|
from .configs import BaseStrategyConfig, DAggerStrategyConfig, RolloutConfig
|
||||||
@@ -109,6 +110,26 @@ def _resolve_action_key_order(
|
|||||||
return policy_action_names
|
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
|
# Sub-contexts
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -431,10 +452,22 @@ def build_rollout_context(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
if isinstance(cfg.inference, SyncInferenceConfig) and any(
|
relative_action_step = next(
|
||||||
isinstance(step, RelativeActionsProcessorStep) and step.enabled
|
(
|
||||||
for step in getattr(preprocessor, "steps", ())
|
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(
|
raise NotImplementedError(
|
||||||
"SyncInferenceEngine does not support policies with relative actions for now."
|
"SyncInferenceEngine does not support policies with relative actions for now."
|
||||||
"Use --inference.type=rtc or remove relative action processor steps from the policy pipeline."
|
"Use --inference.type=rtc or remove relative action processor steps from the policy pipeline."
|
||||||
|
|||||||
@@ -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)
|
_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():
|
def test_sentry_config_defaults():
|
||||||
from lerobot.rollout import SentryStrategyConfig
|
from lerobot.rollout import SentryStrategyConfig
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user