diff --git a/src/lerobot/datasets/compute_stats.py b/src/lerobot/datasets/compute_stats.py index 02ecd81a4..773122fce 100644 --- a/src/lerobot/datasets/compute_stats.py +++ b/src/lerobot/datasets/compute_stats.py @@ -682,6 +682,7 @@ def compute_relative_action_stats( chunk_size: int, exclude_joints: list[str] | None = None, num_workers: int = 0, + state_from_action: bool = False, ) -> dict[str, np.ndarray]: """Compute normalization statistics for relative actions over the full dataset. @@ -700,6 +701,9 @@ def compute_relative_action_stats( num_workers: Number of parallel threads for computation. Values ≤1 mean single-threaded. Numpy releases the GIL so threads give real parallelism here. + state_from_action: Use the current absolute action as state. This is + intended for state-less pose datasets where each action row is the + synchronized measured robot pose. Returns: Statistics dict with keys "mean", "std", "min", "max", "q01", …, "q99". @@ -722,7 +726,7 @@ def compute_relative_action_stats( logging.info("Loading action/state data for relative action stats...") all_actions = np.array(hf_dataset[ACTION], dtype=np.float32) - all_states = np.array(hf_dataset[OBS_STATE], dtype=np.float32) + all_states = all_actions if state_from_action else np.array(hf_dataset[OBS_STATE], dtype=np.float32) episode_indices = np.array(hf_dataset["episode_index"]) valid_starts = _get_valid_chunk_starts(episode_indices, chunk_size) @@ -777,3 +781,50 @@ def compute_relative_action_stats( ) return stats + + +def compute_state_history_stats( + hf_dataset, + features: dict, + history_steps: int, + exclude_joints: list[str] | None = None, + relative: bool = False, +) -> dict[str, np.ndarray]: + """Compute stats for flattened state history synthesized from absolute actions. + + History is left-padded with the first action of each episode, matching dataset + boundary padding. When ``relative`` is enabled, every history pose is expressed + relative to its newest pose while excluded dimensions remain absolute. + """ + if history_steps < 1: + raise ValueError("history_steps must be at least 1") + if exclude_joints is None: + exclude_joints = [] + + actions = np.asarray(hf_dataset[ACTION], dtype=np.float32) + episode_indices = np.asarray(hf_dataset["episode_index"]) + sample_indices = np.arange(len(actions)) + episode_starts = np.maximum.accumulate( + np.where( + np.concatenate(([True], episode_indices[1:] != episode_indices[:-1])), + sample_indices, + 0, + ) + ) + offsets = np.arange(-(history_steps - 1), 1) + history_indices = np.maximum(sample_indices[:, None] + offsets[None, :], episode_starts[:, None]) + history = actions[history_indices].copy() + + if relative: + state_dim = actions.shape[-1] + names = features.get(ACTION, {}).get("names") + mask_step = RelativeActionsProcessorStep( + enabled=True, + exclude_joints=exclude_joints, + action_names=names, + ) + mask = np.asarray(mask_step._build_mask(state_dim), dtype=np.float32) + history[..., : len(mask)] -= history[:, -1:, : len(mask)] * mask[None, None, :] + + flattened = history.reshape(len(history), -1) + return get_feature_stats(flattened, axis=0, keepdims=False) diff --git a/src/lerobot/datasets/dataset_tools.py b/src/lerobot/datasets/dataset_tools.py index 31e075d7c..088de7952 100644 --- a/src/lerobot/datasets/dataset_tools.py +++ b/src/lerobot/datasets/dataset_tools.py @@ -54,6 +54,7 @@ from .compute_stats import ( aggregate_stats, compute_episode_stats, compute_relative_action_stats, + compute_state_history_stats, ) from .dataset_metadata import LeRobotDatasetMetadata from .image_writer import write_image @@ -1566,6 +1567,10 @@ def recompute_stats( relative_exclude_joints: list[str] | None = None, chunk_size: int = 50, num_workers: int = 0, + state_from_action: bool = False, + state_history_steps: int = 1, + relative_state_history: bool = False, + relative_state_exclude_joints: list[str] | None = None, ) -> LeRobotDataset: """Recompute stats.json from scratch by iterating all episodes. @@ -1583,6 +1588,12 @@ def recompute_stats( ``policy.chunk_size``. Only used when ``relative_action=True``. num_workers: Number of parallel threads for relative action stats computation. Values ≤1 mean single-threaded. Only used when ``relative_action=True``. + state_from_action: Use absolute action rows as synthetic state while + computing relative-action stats, and write their absolute statistics + under ``observation.state``. + state_history_steps: Number of consecutive synthesized state samples. + relative_state_history: Express state history relative to its newest pose. + relative_state_exclude_joints: State dimensions to retain as absolute. Returns: The same dataset with updated stats. @@ -1606,7 +1617,19 @@ def recompute_stats( # (matching what the model sees during training) and skip action in the # per-episode pass below. relative_action_stats = None - if relative_action and ACTION in features and OBS_STATE in features: + synthetic_state_stats = None + if state_from_action: + if ACTION not in features: + raise ValueError("state_from_action requires an action feature") + synthetic_state_stats = compute_state_history_stats( + dataset.hf_dataset, + features, + history_steps=state_history_steps, + exclude_joints=relative_state_exclude_joints, + relative=relative_state_history, + ) + + if relative_action and ACTION in features and (OBS_STATE in features or state_from_action): if relative_exclude_joints is None: relative_exclude_joints = ["gripper"] relative_action_stats = compute_relative_action_stats( @@ -1615,6 +1638,7 @@ def recompute_stats( chunk_size=chunk_size, exclude_joints=relative_exclude_joints, num_workers=num_workers, + state_from_action=state_from_action, ) features_to_compute.pop(ACTION, None) @@ -1654,6 +1678,8 @@ def recompute_stats( if relative_action_stats is not None: new_stats[ACTION] = relative_action_stats + if synthetic_state_stats is not None: + new_stats[OBS_STATE] = synthetic_state_stats # Merge: keep existing stats for features we didn't recompute if dataset.meta.stats: diff --git a/src/lerobot/policies/pi05/configuration_pi05.py b/src/lerobot/policies/pi05/configuration_pi05.py index 124e85cc9..5cf53abda 100644 --- a/src/lerobot/policies/pi05/configuration_pi05.py +++ b/src/lerobot/policies/pi05/configuration_pi05.py @@ -56,6 +56,14 @@ class PI05Config(PreTrainedConfig): # Populated at runtime from dataset metadata by make_policy. action_feature_names: list[str] | None = None + # Build proprioception from absolute action samples when the dataset has no + # observation.state. With history_steps=2, training samples request t-1 as + # well as the normal t..t+chunk_size-1 action targets. + state_from_action: bool = False + proprioception_history_steps: int = 1 + use_relative_state_history: bool = False + relative_state_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"]) + # Real-Time Chunking (RTC) configuration rtc_config: RTCConfig | None = None @@ -121,6 +129,9 @@ class PI05Config(PreTrainedConfig): if self.dtype not in ["bfloat16", "float32"]: raise ValueError(f"Invalid dtype: {self.dtype}") + if self.proprioception_history_steps < 1: + raise ValueError("proprioception_history_steps must be at least 1") + def validate_features(self) -> None: """Validate and set up input/output features.""" for i in range(self.empty_cameras): @@ -131,13 +142,6 @@ class PI05Config(PreTrainedConfig): ) self.input_features[key] = empty_camera - if OBS_STATE not in self.input_features: - state_feature = PolicyFeature( - type=FeatureType.STATE, - shape=(self.max_state_dim,), # Padded to max_state_dim - ) - self.input_features[OBS_STATE] = state_feature - if ACTION not in self.output_features: action_feature = PolicyFeature( type=FeatureType.ACTION, @@ -145,6 +149,25 @@ class PI05Config(PreTrainedConfig): ) self.output_features[ACTION] = action_feature + if OBS_STATE not in self.input_features: + state_shape = (self.max_state_dim,) + if self.state_from_action and ACTION in self.output_features: + state_shape = self.output_features[ACTION].shape + state_feature = PolicyFeature( + type=FeatureType.STATE, + shape=state_shape, + ) + self.input_features[OBS_STATE] = state_feature + + state_dim = self.input_features[OBS_STATE].shape[-1] + history_state_dim = state_dim * self.proprioception_history_steps + if history_state_dim > self.max_state_dim: + raise ValueError( + "Flattened proprioception history exceeds max_state_dim: " + f"{state_dim} * {self.proprioception_history_steps} = {history_state_dim} > " + f"{self.max_state_dim}" + ) + def get_optimizer_preset(self) -> AdamWConfig: return AdamWConfig( lr=self.optimizer_lr, @@ -168,7 +191,8 @@ class PI05Config(PreTrainedConfig): @property def action_delta_indices(self) -> list: - return list(range(self.chunk_size)) + history_prefix = self.proprioception_history_steps - 1 if self.state_from_action else 0 + return list(range(-history_prefix, self.chunk_size)) @property def reward_delta_indices(self) -> None: diff --git a/src/lerobot/policies/pi05/processor_pi05.py b/src/lerobot/policies/pi05/processor_pi05.py index 2d015b24f..233bb5426 100644 --- a/src/lerobot/policies/pi05/processor_pi05.py +++ b/src/lerobot/policies/pi05/processor_pi05.py @@ -15,7 +15,7 @@ # limitations under the License. from copy import deepcopy -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Any import numpy as np @@ -48,6 +48,144 @@ from lerobot.utils.constants import ( from .configuration_pi05 import PI05Config +@ProcessorStepRegistry.register(name="pi05_state_from_action_processor_step") +@dataclass +class Pi05StateFromActionProcessorStep(ProcessorStep): + """Synthesize proprioception from absolute actions in state-less datasets. + + The dataset loader supplies ``history_steps - 1`` actions before the normal + target chunk. Those leading samples and action(t) become state history; only + the leading samples are then removed from the action targets. + """ + + enabled: bool = False + history_steps: int = 1 + _inference_history: torch.Tensor | None = field(default=None, init=False, repr=False) + + def __call__(self, transition: EnvTransition) -> EnvTransition: + if not self.enabled: + return transition + + observation = transition.get(TransitionKey.OBSERVATION, {}) + observed_state = observation.get(OBS_STATE) + if observed_state is not None: + # At inference the robot normally provides only the current state and + # there is no action target. Build a rolling history in the processor. + if transition.get(TransitionKey.ACTION) is None and observed_state.ndim == 2: + if self._inference_history is None: + self._inference_history = observed_state.unsqueeze(1).repeat(1, self.history_steps, 1) + else: + self._inference_history = torch.cat( + [self._inference_history[:, 1:], observed_state.unsqueeze(1)], dim=1 + ) + new_transition = transition.copy() + new_observation = dict(observation) + new_observation[OBS_STATE] = self._inference_history.clone() + new_transition[TransitionKey.OBSERVATION] = new_observation + return new_transition + return transition + + action = transition.get(TransitionKey.ACTION) + if action is None: + raise ValueError("Cannot synthesize PI0.5 state without action") + if action.ndim != 3: + raise ValueError(f"Expected batched action chunks with shape (B, T, D), got {action.shape}") + if action.shape[1] < self.history_steps: + raise ValueError( + f"Action chunk has {action.shape[1]} steps, fewer than history_steps={self.history_steps}" + ) + + new_transition = transition.copy() + new_observation = dict(observation) + state = action[:, : self.history_steps].clone() + if self.history_steps == 1: + state = state[:, 0] + new_observation[OBS_STATE] = state + new_transition[TransitionKey.OBSERVATION] = new_observation + new_transition[TransitionKey.ACTION] = action[:, self.history_steps - 1 :] + return new_transition + + def get_config(self) -> dict[str, Any]: + return {"enabled": self.enabled, "history_steps": self.history_steps} + + def reset(self) -> None: + self._inference_history = None + + def transform_features( + self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] + ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: + return features + + +@ProcessorStepRegistry.register(name="pi05_flatten_state_history_processor_step") +@dataclass +class Pi05FlattenStateHistoryProcessorStep(ProcessorStep): + """Optionally relativize raw state history, then flatten it for PI0.5.""" + + history_steps: int = 1 + max_state_dim: int = 32 + relative: bool = False + exclude_joints: list[str] = field(default_factory=list) + state_names: list[str] | None = None + + def __call__(self, transition: EnvTransition) -> EnvTransition: + observation = transition.get(TransitionKey.OBSERVATION, {}) + state = observation.get(OBS_STATE) + if state is None: + raise ValueError("State is required for PI05") + if self.history_steps == 1 and state.ndim == 2: + state = state.unsqueeze(1) + if state.ndim != 3 or state.shape[1] != self.history_steps: + raise ValueError( + f"Expected state history with shape (B, {self.history_steps}, D), got {state.shape}" + ) + + flattened_dim = state.shape[1] * state.shape[2] + if flattened_dim > self.max_state_dim: + raise ValueError( + f"Flattened state history has {flattened_dim} dimensions, above max_state_dim={self.max_state_dim}" + ) + + processed_state = state.clone() + if self.relative: + mask_step = RelativeActionsProcessorStep( + enabled=True, + exclude_joints=self.exclude_joints, + action_names=self.state_names, + ) + mask = torch.tensor( + mask_step._build_mask(state.shape[-1]), dtype=state.dtype, device=state.device + ) + reference = state[:, -1:, : mask.shape[0]] + processed_state[..., : mask.shape[0]] -= reference * mask + + new_transition = transition.copy() + new_observation = dict(observation) + new_observation[OBS_STATE] = processed_state.flatten(start_dim=1) + new_transition[TransitionKey.OBSERVATION] = new_observation + return new_transition + + def get_config(self) -> dict[str, Any]: + return { + "history_steps": self.history_steps, + "max_state_dim": self.max_state_dim, + "relative": self.relative, + "exclude_joints": self.exclude_joints, + "state_names": self.state_names, + } + + def transform_features( + self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]] + ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]: + transformed = deepcopy(features) + for feature_group in transformed.values(): + state_feature = feature_group.get(OBS_STATE) + if state_feature is not None and self.history_steps > 1: + state_dim = state_feature.shape[-1] * self.history_steps + feature_group[OBS_STATE] = PolicyFeature(type=state_feature.type, shape=(state_dim,)) + return transformed + + @ProcessorStepRegistry.register(name="pi05_prepare_state_tokenizer_processor_step") @dataclass class Pi05PrepareStateTokenizerProcessorStep(ProcessorStep): @@ -139,7 +277,18 @@ def make_pi05_pre_post_processors( input_steps: list[ProcessorStep] = [ RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one AddBatchDimensionProcessorStep(), + Pi05StateFromActionProcessorStep( + enabled=config.state_from_action, + history_steps=config.proprioception_history_steps, + ), relative_step, + Pi05FlattenStateHistoryProcessorStep( + history_steps=config.proprioception_history_steps, + max_state_dim=config.max_state_dim, + relative=config.use_relative_state_history, + exclude_joints=config.relative_state_exclude_joints, + state_names=config.action_feature_names, + ), # NOTE: NormalizerProcessorStep MUST come before Pi05PrepareStateTokenizerProcessorStep # because the tokenizer step expects normalized state in [-1, 1] range for discretization NormalizerProcessorStep( diff --git a/src/lerobot/processor/relative_action_processor.py b/src/lerobot/processor/relative_action_processor.py index 5b1039291..3f5873463 100644 --- a/src/lerobot/processor/relative_action_processor.py +++ b/src/lerobot/processor/relative_action_processor.py @@ -126,20 +126,24 @@ class RelativeActionsProcessorStep(ProcessorStep): observation = transition.get(TransitionKey.OBSERVATION, {}) state = observation.get(OBS_STATE) if observation else None + # State history has shape (B, H, D). Relative actions are referenced to + # the newest proprioceptive state, not the whole history tensor. + reference_state = state[:, -1] if state is not None and state.ndim == 3 else state + # Always cache state for the paired AbsoluteActionsProcessorStep - if state is not None: - self._last_state = state + if reference_state is not None: + self._last_state = reference_state if not self.enabled: return transition new_transition = transition.copy() action = new_transition.get(TransitionKey.ACTION) - if action is None or state is None: + if action is None or reference_state is None: return new_transition mask = self._build_mask(action.shape[-1]) - new_transition[TransitionKey.ACTION] = to_relative_actions(action, state, mask) + new_transition[TransitionKey.ACTION] = to_relative_actions(action, reference_state, mask) return new_transition def get_cached_state(self) -> torch.Tensor | None: diff --git a/tests/processor/test_pi05_state_from_action.py b/tests/processor/test_pi05_state_from_action.py new file mode 100644 index 000000000..3a85a9395 --- /dev/null +++ b/tests/processor/test_pi05_state_from_action.py @@ -0,0 +1,175 @@ +#!/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. + +import numpy as np +import pytest +import torch + +pytest.importorskip("transformers") + +from lerobot.datasets.compute_stats import ( # noqa: E402 + compute_relative_action_stats, + compute_state_history_stats, +) +from lerobot.policies.pi05.configuration_pi05 import PI05Config # noqa: E402 +from lerobot.policies.pi05.processor_pi05 import ( # noqa: E402 + Pi05FlattenStateHistoryProcessorStep, + Pi05StateFromActionProcessorStep, +) +from lerobot.processor.relative_action_processor import ( # noqa: E402 + AbsoluteActionsProcessorStep, + RelativeActionsProcessorStep, +) +from lerobot.types import TransitionKey # noqa: E402 +from lerobot.utils.constants import OBS_STATE # noqa: E402 + + +def _transition(action: torch.Tensor | None, state: torch.Tensor | None = None) -> dict: + observation = {} if state is None else {OBS_STATE: state} + return { + TransitionKey.OBSERVATION: observation, + TransitionKey.ACTION: action, + TransitionKey.REWARD: None, + TransitionKey.DONE: None, + TransitionKey.TRUNCATED: None, + TransitionKey.COMPLEMENTARY_DATA: {}, + } + + +def test_pi05_config_requests_action_history_prefix(): + config = PI05Config( + device="cpu", + chunk_size=4, + n_action_steps=4, + state_from_action=True, + proprioception_history_steps=2, + ) + + assert config.action_delta_indices == [-1, 0, 1, 2, 3] + + +def test_state_from_action_extracts_history_and_preserves_target_horizon(): + action = torch.arange(2 * 5 * 3, dtype=torch.float32).reshape(2, 5, 3) + step = Pi05StateFromActionProcessorStep(enabled=True, history_steps=2) + + result = step(_transition(action)) + + torch.testing.assert_close(result[TransitionKey.OBSERVATION][OBS_STATE], action[:, :2]) + torch.testing.assert_close(result[TransitionKey.ACTION], action[:, 1:]) + + +def test_relative_actions_use_newest_state_in_history_and_roundtrip(): + state_history = torch.tensor([[[1.0, 10.0], [2.0, 20.0]]]) + absolute = torch.tensor([[[3.0, 30.0], [4.0, 40.0]]]) + relative_step = RelativeActionsProcessorStep(enabled=True) + absolute_step = AbsoluteActionsProcessorStep(enabled=True, relative_step=relative_step) + + relative = relative_step(_transition(absolute, state_history)) + expected = torch.tensor([[[1.0, 10.0], [2.0, 20.0]]]) + torch.testing.assert_close(relative[TransitionKey.ACTION], expected) + + recovered = absolute_step(_transition(relative[TransitionKey.ACTION])) + torch.testing.assert_close(recovered[TransitionKey.ACTION], absolute) + + +def test_flatten_state_history_preserves_chronological_order(): + state_history = torch.tensor([[[1.0, 2.0], [3.0, 4.0]]]) + step = Pi05FlattenStateHistoryProcessorStep(history_steps=2, max_state_dim=4) + + result = step(_transition(torch.zeros(1, 2, 2), state_history)) + + torch.testing.assert_close( + result[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[1.0, 2.0, 3.0, 4.0]]) + ) + + +def test_state_history_can_be_relative_with_absolute_gripper(): + state_history = torch.tensor([[[1.0, 10.0, 0.2], [3.0, 20.0, 0.4]]]) + step = Pi05FlattenStateHistoryProcessorStep( + history_steps=2, + max_state_dim=6, + relative=True, + exclude_joints=["gripper"], + state_names=["x", "y", "gripper_width"], + ) + + result = step(_transition(torch.zeros(1, 2, 3), state_history)) + + torch.testing.assert_close( + result[TransitionKey.OBSERVATION][OBS_STATE], + torch.tensor([[-2.0, -10.0, 0.2, 0.0, 0.0, 0.4]]), + ) + + +def test_inference_state_history_is_rolled_and_reset(): + step = Pi05StateFromActionProcessorStep(enabled=True, history_steps=2) + + first = step(_transition(None, torch.tensor([[1.0, 2.0]]))) + second = step(_transition(None, torch.tensor([[3.0, 4.0]]))) + torch.testing.assert_close( + first[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[[1.0, 2.0], [1.0, 2.0]]]) + ) + torch.testing.assert_close( + second[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[[1.0, 2.0], [3.0, 4.0]]]) + ) + + step.reset() + reset = step(_transition(None, torch.tensor([[5.0, 6.0]]))) + torch.testing.assert_close( + reset[TransitionKey.OBSERVATION][OBS_STATE], torch.tensor([[[5.0, 6.0], [5.0, 6.0]]]) + ) + + +def test_flatten_state_history_checks_max_state_dim(): + step = Pi05FlattenStateHistoryProcessorStep(history_steps=2, max_state_dim=3) + + with pytest.raises(ValueError, match="above max_state_dim"): + step(_transition(torch.zeros(1, 2, 2), torch.zeros(1, 2, 2))) + + +def test_relative_stats_can_use_absolute_action_as_state(): + actions = np.asarray([[0.0, 0.0], [1.0, 2.0], [2.0, 4.0], [3.0, 6.0]], dtype=np.float32) + dataset = {"action": actions, "episode_index": np.zeros(4, dtype=np.int64)} + features = {"action": {"shape": [2], "names": ["x", "y"]}} + + stats = compute_relative_action_stats( + dataset, + features, + chunk_size=2, + state_from_action=True, + ) + + np.testing.assert_allclose(stats["mean"], [0.5, 1.0]) + + +def test_relative_state_history_stats_match_processor_representation(): + actions = np.asarray( + [[0.0, 0.1], [1.0, 0.2], [3.0, 0.3]], + dtype=np.float32, + ) + dataset = {"action": actions, "episode_index": np.zeros(3, dtype=np.int64)} + features = {"action": {"shape": [2], "names": ["x", "gripper_width"]}} + + stats = compute_state_history_stats( + dataset, + features, + history_steps=2, + exclude_joints=["gripper"], + relative=True, + ) + + expected = np.asarray([[0.0, 0.1, 0.0, 0.1], [-1.0, 0.1, 0.0, 0.2], [-2.0, 0.2, 0.0, 0.3]]) + np.testing.assert_allclose(stats["mean"], expected.mean(axis=0))