mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-23 09:46:00 +00:00
feat(pi05): synthesize relative proprioceptive history
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
Reference in New Issue
Block a user