feat(pi05): compose relative poses in SE3

This commit is contained in:
Pepijn
2026-07-21 20:43:04 +02:00
parent 7e1077f19a
commit 8e12a5351a
9 changed files with 392 additions and 27 deletions
+39 -4
View File
@@ -18,8 +18,9 @@ from __future__ import annotations
import logging import logging
import numpy as np import numpy as np
import torch
from lerobot.processor import RelativeActionsProcessorStep from lerobot.processor import RelativeActionsProcessorStep, to_relative_actions
from lerobot.utils.constants import ACTION, OBS_STATE from lerobot.utils.constants import ACTION, OBS_STATE
from .io_utils import load_image_as_numpy from .io_utils import load_image_as_numpy
@@ -660,6 +661,8 @@ def _compute_relative_chunk_batch(
all_states: np.ndarray, all_states: np.ndarray,
chunk_size: int, chunk_size: int,
relative_mask: np.ndarray, relative_mask: np.ndarray,
pose_representation: str = "componentwise",
se3_pose_groups: list[list[int]] | None = None,
) -> np.ndarray: ) -> np.ndarray:
"""Vectorised relative-action computation for a batch of start indices. """Vectorised relative-action computation for a batch of start indices.
@@ -671,6 +674,18 @@ def _compute_relative_chunk_batch(
frame_idx = start_indices[:, None] + offsets[None, :] frame_idx = start_indices[:, None] + offsets[None, :]
chunks = all_actions[frame_idx].copy() chunks = all_actions[frame_idx].copy()
states = all_states[start_indices] states = all_states[start_indices]
if pose_representation == "se3":
return (
to_relative_actions(
torch.from_numpy(chunks),
torch.from_numpy(states),
relative_mask.astype(bool).tolist(),
pose_representation=pose_representation,
se3_pose_groups=se3_pose_groups,
)
.numpy()
.reshape(-1, all_actions.shape[1])
)
mask_dim = len(relative_mask) mask_dim = len(relative_mask)
chunks[:, :, :mask_dim] -= states[:, None, :mask_dim] * relative_mask[None, None, :] chunks[:, :, :mask_dim] -= states[:, None, :mask_dim] * relative_mask[None, None, :]
return chunks.reshape(-1, all_actions.shape[1]) return chunks.reshape(-1, all_actions.shape[1])
@@ -683,6 +698,8 @@ def compute_relative_action_stats(
exclude_joints: list[str] | None = None, exclude_joints: list[str] | None = None,
num_workers: int = 0, num_workers: int = 0,
state_from_action: bool = False, state_from_action: bool = False,
pose_representation: str = "componentwise",
se3_pose_groups: list[list[int]] | None = None,
) -> dict[str, np.ndarray]: ) -> dict[str, np.ndarray]:
"""Compute normalization statistics for relative actions over the full dataset. """Compute normalization statistics for relative actions over the full dataset.
@@ -758,6 +775,8 @@ def compute_relative_action_stats(
all_states, all_states,
chunk_size, chunk_size,
relative_mask, relative_mask,
pose_representation,
se3_pose_groups,
) )
for batch in batches for batch in batches
] ]
@@ -766,7 +785,15 @@ def compute_relative_action_stats(
else: else:
for batch in batches: for batch in batches:
running_stats.update( running_stats.update(
_compute_relative_chunk_batch(batch, all_actions, all_states, chunk_size, relative_mask) _compute_relative_chunk_batch(
batch,
all_actions,
all_states,
chunk_size,
relative_mask,
pose_representation,
se3_pose_groups,
)
) )
stats = running_stats.get_statistics() stats = running_stats.get_statistics()
@@ -789,6 +816,8 @@ def compute_state_history_stats(
history_steps: int, history_steps: int,
exclude_joints: list[str] | None = None, exclude_joints: list[str] | None = None,
relative: bool = False, relative: bool = False,
pose_representation: str = "componentwise",
se3_pose_groups: list[list[int]] | None = None,
) -> dict[str, np.ndarray]: ) -> dict[str, np.ndarray]:
"""Compute stats for flattened state history synthesized from absolute actions. """Compute stats for flattened state history synthesized from absolute actions.
@@ -823,8 +852,14 @@ def compute_state_history_stats(
exclude_joints=exclude_joints, exclude_joints=exclude_joints,
action_names=names, action_names=names,
) )
mask = np.asarray(mask_step._build_mask(state_dim), dtype=np.float32) mask = mask_step._build_mask(state_dim)
history[..., : len(mask)] -= history[:, -1:, : len(mask)] * mask[None, None, :] history = to_relative_actions(
torch.from_numpy(history),
torch.from_numpy(history[:, -1].copy()),
mask,
pose_representation=pose_representation,
se3_pose_groups=se3_pose_groups,
).numpy()
flattened = history.reshape(len(history), -1) flattened = history.reshape(len(history), -1)
return get_feature_stats(flattened, axis=0, keepdims=False) return get_feature_stats(flattened, axis=0, keepdims=False)
+9
View File
@@ -1571,6 +1571,8 @@ def recompute_stats(
state_history_steps: int = 1, state_history_steps: int = 1,
relative_state_history: bool = False, relative_state_history: bool = False,
relative_state_exclude_joints: list[str] | None = None, relative_state_exclude_joints: list[str] | None = None,
relative_pose_representation: str = "componentwise",
relative_se3_pose_groups: list[list[int]] | None = None,
) -> LeRobotDataset: ) -> LeRobotDataset:
"""Recompute stats.json from scratch by iterating all episodes. """Recompute stats.json from scratch by iterating all episodes.
@@ -1594,6 +1596,9 @@ def recompute_stats(
state_history_steps: Number of consecutive synthesized state samples. state_history_steps: Number of consecutive synthesized state samples.
relative_state_history: Express state history relative to its newest pose. relative_state_history: Express state history relative to its newest pose.
relative_state_exclude_joints: State dimensions to retain as absolute. relative_state_exclude_joints: State dimensions to retain as absolute.
relative_pose_representation: ``componentwise`` for legacy subtraction or
``se3`` for ``inv(T_current) @ T_target`` pose composition.
relative_se3_pose_groups: Six-index xyz+rotation-vector pose groups.
Returns: Returns:
The same dataset with updated stats. The same dataset with updated stats.
@@ -1627,6 +1632,8 @@ def recompute_stats(
history_steps=state_history_steps, history_steps=state_history_steps,
exclude_joints=relative_state_exclude_joints, exclude_joints=relative_state_exclude_joints,
relative=relative_state_history, relative=relative_state_history,
pose_representation=relative_pose_representation,
se3_pose_groups=relative_se3_pose_groups,
) )
if relative_action and ACTION in features and (OBS_STATE in features or state_from_action): if relative_action and ACTION in features and (OBS_STATE in features or state_from_action):
@@ -1639,6 +1646,8 @@ def recompute_stats(
exclude_joints=relative_exclude_joints, exclude_joints=relative_exclude_joints,
num_workers=num_workers, num_workers=num_workers,
state_from_action=state_from_action, state_from_action=state_from_action,
pose_representation=relative_pose_representation,
se3_pose_groups=relative_se3_pose_groups,
) )
features_to_compute.pop(ACTION, None) features_to_compute.pop(ACTION, None)
@@ -55,6 +55,10 @@ class PI05Config(PreTrainedConfig):
relative_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"]) relative_exclude_joints: list[str] = field(default_factory=lambda: ["gripper"])
# Populated at runtime from dataset metadata by make_policy. # Populated at runtime from dataset metadata by make_policy.
action_feature_names: list[str] | None = None action_feature_names: list[str] | None = None
# ``se3`` uses inv(T_current) @ T_target for each xyz+rotation-vector pose group.
# ``componentwise`` preserves the legacy action - state behavior.
relative_pose_representation: str = "componentwise"
relative_se3_pose_groups: list[list[int]] = field(default_factory=lambda: [list(range(6))])
# Build proprioception from absolute action samples when the dataset has no # Build proprioception from absolute action samples when the dataset has no
# observation.state. With history_steps=2, training samples request t-1 as # observation.state. With history_steps=2, training samples request t-1 as
@@ -132,6 +136,17 @@ class PI05Config(PreTrainedConfig):
if self.proprioception_history_steps < 1: if self.proprioception_history_steps < 1:
raise ValueError("proprioception_history_steps must be at least 1") raise ValueError("proprioception_history_steps must be at least 1")
if self.relative_pose_representation not in {"componentwise", "se3"}:
raise ValueError(
"relative_pose_representation must be either 'componentwise' or 'se3', got "
f"{self.relative_pose_representation!r}"
)
for group in self.relative_se3_pose_groups:
if len(group) != 6 or len(set(group)) != 6 or any(index < 0 for index in group):
raise ValueError(f"Invalid six-index SE(3) pose group: {group}")
if self.relative_pose_representation == "se3" and not self.relative_se3_pose_groups:
raise ValueError("relative_pose_representation='se3' requires relative_se3_pose_groups")
def validate_features(self) -> None: def validate_features(self) -> None:
"""Validate and set up input/output features.""" """Validate and set up input/output features."""
for i in range(self.empty_cameras): for i in range(self.empty_cameras):
+15 -4
View File
@@ -36,6 +36,7 @@ from lerobot.processor import (
TokenizerProcessorStep, TokenizerProcessorStep,
UnnormalizerProcessorStep, UnnormalizerProcessorStep,
policy_action_to_transition, policy_action_to_transition,
to_relative_actions,
transition_to_policy_action, transition_to_policy_action,
) )
from lerobot.types import EnvTransition, TransitionKey from lerobot.types import EnvTransition, TransitionKey
@@ -127,6 +128,8 @@ class Pi05FlattenStateHistoryProcessorStep(ProcessorStep):
relative: bool = False relative: bool = False
exclude_joints: list[str] = field(default_factory=list) exclude_joints: list[str] = field(default_factory=list)
state_names: list[str] | None = None state_names: list[str] | None = None
pose_representation: str = "componentwise"
se3_pose_groups: list[list[int]] = field(default_factory=list)
def __call__(self, transition: EnvTransition) -> EnvTransition: def __call__(self, transition: EnvTransition) -> EnvTransition:
observation = transition.get(TransitionKey.OBSERVATION, {}) observation = transition.get(TransitionKey.OBSERVATION, {})
@@ -153,11 +156,13 @@ class Pi05FlattenStateHistoryProcessorStep(ProcessorStep):
exclude_joints=self.exclude_joints, exclude_joints=self.exclude_joints,
action_names=self.state_names, action_names=self.state_names,
) )
mask = torch.tensor( processed_state = to_relative_actions(
mask_step._build_mask(state.shape[-1]), dtype=state.dtype, device=state.device state,
state[:, -1],
mask_step._build_mask(state.shape[-1]),
pose_representation=self.pose_representation,
se3_pose_groups=self.se3_pose_groups,
) )
reference = state[:, -1:, : mask.shape[0]]
processed_state[..., : mask.shape[0]] -= reference * mask
new_transition = transition.copy() new_transition = transition.copy()
new_observation = dict(observation) new_observation = dict(observation)
@@ -172,6 +177,8 @@ class Pi05FlattenStateHistoryProcessorStep(ProcessorStep):
"relative": self.relative, "relative": self.relative,
"exclude_joints": self.exclude_joints, "exclude_joints": self.exclude_joints,
"state_names": self.state_names, "state_names": self.state_names,
"pose_representation": self.pose_representation,
"se3_pose_groups": self.se3_pose_groups,
} }
def transform_features( def transform_features(
@@ -271,6 +278,8 @@ def make_pi05_pre_post_processors(
enabled=config.use_relative_actions, enabled=config.use_relative_actions,
exclude_joints=getattr(config, "relative_exclude_joints", []), exclude_joints=getattr(config, "relative_exclude_joints", []),
action_names=getattr(config, "action_feature_names", None), action_names=getattr(config, "action_feature_names", None),
pose_representation=config.relative_pose_representation,
se3_pose_groups=config.relative_se3_pose_groups,
) )
# OpenPI order: raw → relative → normalize → model → unnormalize → absolute # OpenPI order: raw → relative → normalize → model → unnormalize → absolute
@@ -288,6 +297,8 @@ def make_pi05_pre_post_processors(
relative=config.use_relative_state_history, relative=config.use_relative_state_history,
exclude_joints=config.relative_state_exclude_joints, exclude_joints=config.relative_state_exclude_joints,
state_names=config.action_feature_names, state_names=config.action_feature_names,
pose_representation=config.relative_pose_representation,
se3_pose_groups=config.relative_se3_pose_groups,
), ),
# NOTE: NormalizerProcessorStep MUST come before Pi05PrepareStateTokenizerProcessorStep # NOTE: NormalizerProcessorStep MUST come before Pi05PrepareStateTokenizerProcessorStep
# because the tokenizer step expects normalized state in [-1, 1] range for discretization # because the tokenizer step expects normalized state in [-1, 1] range for discretization
+6 -2
View File
@@ -90,7 +90,9 @@ from .relative_action_processor import (
AbsoluteActionsProcessorStep, AbsoluteActionsProcessorStep,
RelativeActionsProcessorStep, RelativeActionsProcessorStep,
to_absolute_actions, to_absolute_actions,
to_absolute_se3_pose,
to_relative_actions, to_relative_actions,
to_relative_se3_pose,
) )
from .rename_processor import RenameObservationsProcessorStep, rename_stats from .rename_processor import RenameObservationsProcessorStep, rename_stats
from .tokenizer_processor import ActionTokenizerProcessorStep, TokenizerProcessorStep from .tokenizer_processor import ActionTokenizerProcessorStep, TokenizerProcessorStep
@@ -135,6 +137,10 @@ __all__ = [
"make_default_robot_observation_processor", "make_default_robot_observation_processor",
"AbsoluteActionsProcessorStep", "AbsoluteActionsProcessorStep",
"RelativeActionsProcessorStep", "RelativeActionsProcessorStep",
"to_absolute_actions",
"to_absolute_se3_pose",
"to_relative_actions",
"to_relative_se3_pose",
"MapDeltaActionToRobotActionStep", "MapDeltaActionToRobotActionStep",
"MapTensorToDeltaActionDictStep", "MapTensorToDeltaActionDictStep",
"NewLineTaskProcessorStep", "NewLineTaskProcessorStep",
@@ -168,8 +174,6 @@ __all__ = [
"transition_to_batch", "transition_to_batch",
"TransitionKey", "TransitionKey",
"TruncatedProcessorStep", "TruncatedProcessorStep",
"to_absolute_actions",
"to_relative_actions",
"UnnormalizerProcessorStep", "UnnormalizerProcessorStep",
"VanillaObservationProcessorStep", "VanillaObservationProcessorStep",
] ]
@@ -34,57 +34,206 @@ __all__ = [
"AbsoluteActionsProcessorStep", "AbsoluteActionsProcessorStep",
"to_relative_actions", "to_relative_actions",
"to_absolute_actions", "to_absolute_actions",
"to_relative_se3_pose",
"to_absolute_se3_pose",
] ]
def to_relative_actions(actions: Tensor, state: Tensor, mask: Sequence[bool]) -> Tensor: def _rotvec_to_quaternion(rotvec: Tensor) -> Tensor:
"""Convert absolute actions to relative: relative = action - state (for masked dims). angle = torch.linalg.vector_norm(rotvec, dim=-1, keepdim=True)
angle_sq = angle.square()
small_scale = 0.5 - angle_sq / 48.0 + angle_sq.square() / 3840.0
scale = torch.where(angle > 1e-6, torch.sin(angle / 2.0) / angle.clamp_min(1e-12), small_scale)
return torch.cat((torch.cos(angle / 2.0), rotvec * scale), dim=-1)
def _quaternion_to_rotvec(quaternion: Tensor) -> Tensor:
quaternion = quaternion / torch.linalg.vector_norm(quaternion, dim=-1, keepdim=True).clamp_min(1e-12)
quaternion = quaternion * torch.where(quaternion[..., :1] < 0, -1.0, 1.0)
vector = quaternion[..., 1:]
sin_half_angle = torch.linalg.vector_norm(vector, dim=-1, keepdim=True)
angle = 2.0 * torch.atan2(sin_half_angle, quaternion[..., :1].clamp_min(0.0))
small_scale = 2.0 + sin_half_angle.square() / 3.0
scale = torch.where(
sin_half_angle > 1e-6,
angle / sin_half_angle.clamp_min(1e-12),
small_scale,
)
return vector * scale
def _quaternion_multiply(left: Tensor, right: Tensor) -> Tensor:
left_w, left_xyz = left[..., :1], left[..., 1:]
right_w, right_xyz = right[..., :1], right[..., 1:]
return torch.cat(
(
left_w * right_w - (left_xyz * right_xyz).sum(dim=-1, keepdim=True),
left_w * right_xyz + right_w * left_xyz + torch.linalg.cross(left_xyz, right_xyz, dim=-1),
),
dim=-1,
)
def _quaternion_conjugate(quaternion: Tensor) -> Tensor:
return torch.cat((quaternion[..., :1], -quaternion[..., 1:]), dim=-1)
def _quaternion_rotate(quaternion: Tensor, vector: Tensor) -> Tensor:
quaternion_xyz = quaternion[..., 1:]
uv = torch.linalg.cross(quaternion_xyz, vector, dim=-1)
uuv = torch.linalg.cross(quaternion_xyz, uv, dim=-1)
return vector + 2.0 * (quaternion[..., :1] * uv + uuv)
def to_relative_se3_pose(target_pose: Tensor, reference_pose: Tensor) -> Tensor:
"""Encode a pose as ``inv(T_reference) @ T_target``.
Poses use ``[x, y, z, rx, ry, rz]`` with an axis-angle rotation vector.
The relative translation is therefore expressed in the reference EE frame.
"""
if target_pose.shape[-1] != 6 or reference_pose.shape[-1] != 6:
raise ValueError("SE(3) poses must have six values: xyz followed by a rotation vector")
reference_quaternion = _rotvec_to_quaternion(reference_pose[..., 3:])
target_quaternion = _rotvec_to_quaternion(target_pose[..., 3:])
inverse_reference_quaternion = _quaternion_conjugate(reference_quaternion)
relative_translation = _quaternion_rotate(
inverse_reference_quaternion, target_pose[..., :3] - reference_pose[..., :3]
)
relative_quaternion = _quaternion_multiply(inverse_reference_quaternion, target_quaternion)
return torch.cat((relative_translation, _quaternion_to_rotvec(relative_quaternion)), dim=-1)
def to_absolute_se3_pose(relative_pose: Tensor, reference_pose: Tensor) -> Tensor:
"""Decode a pose with ``T_target = T_reference @ T_relative``."""
if relative_pose.shape[-1] != 6 or reference_pose.shape[-1] != 6:
raise ValueError("SE(3) poses must have six values: xyz followed by a rotation vector")
reference_quaternion = _rotvec_to_quaternion(reference_pose[..., 3:])
relative_quaternion = _rotvec_to_quaternion(relative_pose[..., 3:])
target_translation = reference_pose[..., :3] + _quaternion_rotate(
reference_quaternion, relative_pose[..., :3]
)
target_quaternion = _quaternion_multiply(reference_quaternion, relative_quaternion)
return torch.cat((target_translation, _quaternion_to_rotvec(target_quaternion)), dim=-1)
def _broadcast_reference(actions: Tensor, state: Tensor) -> Tensor:
if state.device != actions.device or state.dtype != actions.dtype:
state = state.to(device=actions.device, dtype=actions.dtype)
if actions.ndim == state.ndim + 1:
state = state.unsqueeze(-2)
return state
def _validate_se3_pose_groups(
pose_representation: str,
se3_pose_groups: Sequence[Sequence[int]] | None,
mask: Sequence[bool],
action_dim: int,
) -> list[list[int]]:
if pose_representation not in {"componentwise", "se3"}:
raise ValueError(
f"Unsupported pose_representation={pose_representation!r}; expected 'componentwise' or 'se3'"
)
if pose_representation == "componentwise":
return []
if not se3_pose_groups:
raise ValueError("pose_representation='se3' requires at least one six-index se3_pose_group")
normalized_groups: list[list[int]] = []
used_indices: set[int] = set()
for raw_group in se3_pose_groups:
group = [int(index) for index in raw_group]
if len(group) != 6:
raise ValueError(f"Each SE(3) pose group must contain six indices, got {group}")
if len(set(group)) != 6 or any(index < 0 or index >= action_dim for index in group):
raise ValueError(f"Invalid SE(3) pose group for action_dim={action_dim}: {group}")
if any(index >= len(mask) for index in group):
raise ValueError(f"SE(3) pose group lies outside the relative mask: {group}")
if used_indices.intersection(group):
raise ValueError(f"SE(3) pose groups must not overlap: {group}")
group_mask = [bool(mask[index]) for index in group]
if any(group_mask) and not all(group_mask):
raise ValueError(f"An SE(3) pose group must be wholly relative or wholly absolute: {group}")
used_indices.update(group)
if all(group_mask):
normalized_groups.append(group)
return normalized_groups
def to_relative_actions(
actions: Tensor,
state: Tensor,
mask: Sequence[bool],
*,
pose_representation: str = "componentwise",
se3_pose_groups: Sequence[Sequence[int]] | None = None,
) -> Tensor:
"""Convert absolute actions to a configured relative representation.
Component-wise mode computes ``action - state``. SE(3) mode computes
``inv(T_state) @ T_action`` for each configured pose group.
Args: Args:
actions: (B, T, action_dim) or (B, action_dim). actions: (B, T, action_dim) or (B, action_dim).
state: (B, state_dim). Broadcast across time dimension. state: (B, state_dim). Broadcast across time dimension.
mask: Which dims to convert. Can be shorter than action_dim. mask: Which dims to convert. Can be shorter than action_dim.
""" """
groups = _validate_se3_pose_groups(pose_representation, se3_pose_groups, mask, actions.shape[-1])
mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device) mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device)
dims = mask_t.shape[0] dims = mask_t.shape[0]
# Align state to the same device/dtype as actions. _last_state is cached before # Align state to the same device/dtype as actions. _last_state is cached before
# DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA. # DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA.
if state.device != actions.device or state.dtype != actions.dtype: state = _broadcast_reference(actions, state)
state = state.to(device=actions.device, dtype=actions.dtype) component_mask = mask_t.clone()
state_offset = state[..., :dims] * mask_t for group in groups:
if actions.ndim == 3: component_mask[group] = 0
state_offset = state_offset.unsqueeze(-2) state_offset = state[..., :dims] * component_mask
actions = actions.clone() actions = actions.clone()
actions[..., :dims] -= state_offset actions[..., :dims] -= state_offset
for group in groups:
actions[..., group] = to_relative_se3_pose(actions[..., group], state[..., group])
return actions return actions
def to_absolute_actions(actions: Tensor, state: Tensor, mask: Sequence[bool]) -> Tensor: def to_absolute_actions(
"""Convert relative actions back to absolute: absolute = relative + state (for masked dims). actions: Tensor,
state: Tensor,
mask: Sequence[bool],
*,
pose_representation: str = "componentwise",
se3_pose_groups: Sequence[Sequence[int]] | None = None,
) -> Tensor:
"""Convert relative actions back to absolute actions.
Component-wise mode computes ``relative + state``. SE(3) mode computes
``T_state @ T_relative`` for each configured pose group.
Args: Args:
actions: (B, T, action_dim) or (B, action_dim). actions: (B, T, action_dim) or (B, action_dim).
state: (B, state_dim). Broadcast across time dimension. state: (B, state_dim). Broadcast across time dimension.
mask: Which dims to convert. Can be shorter than action_dim. mask: Which dims to convert. Can be shorter than action_dim.
""" """
groups = _validate_se3_pose_groups(pose_representation, se3_pose_groups, mask, actions.shape[-1])
mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device) mask_t = torch.tensor(mask, dtype=actions.dtype, device=actions.device)
dims = mask_t.shape[0] dims = mask_t.shape[0]
# Align state to the same device/dtype as actions. _last_state is cached before # Align state to the same device/dtype as actions. _last_state is cached before
# DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA. # DeviceProcessorStep moves the transition, so it can be on CPU while actions are on CUDA.
if state.device != actions.device or state.dtype != actions.dtype: state = _broadcast_reference(actions, state)
state = state.to(device=actions.device, dtype=actions.dtype) component_mask = mask_t.clone()
state_offset = state[..., :dims] * mask_t for group in groups:
if actions.ndim == 3: component_mask[group] = 0
state_offset = state_offset.unsqueeze(-2) state_offset = state[..., :dims] * component_mask
actions = actions.clone() actions = actions.clone()
actions[..., :dims] += state_offset actions[..., :dims] += state_offset
for group in groups:
actions[..., group] = to_absolute_se3_pose(actions[..., group], state[..., group])
return actions return actions
@ProcessorStepRegistry.register("relative_actions_processor") @ProcessorStepRegistry.register("relative_actions_processor")
@dataclass @dataclass
class RelativeActionsProcessorStep(ProcessorStep): class RelativeActionsProcessorStep(ProcessorStep):
"""Converts absolute actions to relative actions (action -= state) for masked dimensions. """Converts absolute actions to the configured relative representation.
Mirrors OpenPI's DeltaActions transform. Applied during preprocessing so the model Mirrors OpenPI's DeltaActions transform. Applied during preprocessing so the model
trains on relative offsets instead of absolute positions. trains on relative offsets instead of absolute positions.
@@ -101,6 +250,8 @@ class RelativeActionsProcessorStep(ProcessorStep):
enabled: bool = False enabled: bool = False
exclude_joints: list[str] = field(default_factory=list) exclude_joints: list[str] = field(default_factory=list)
action_names: list[str] | None = None action_names: list[str] | None = None
pose_representation: str = "componentwise"
se3_pose_groups: list[list[int]] = field(default_factory=list)
_last_state: torch.Tensor | None = field(default=None, init=False, repr=False) _last_state: torch.Tensor | None = field(default=None, init=False, repr=False)
def _build_mask(self, action_dim: int) -> list[bool]: def _build_mask(self, action_dim: int) -> list[bool]:
@@ -143,7 +294,13 @@ class RelativeActionsProcessorStep(ProcessorStep):
return new_transition return new_transition
mask = self._build_mask(action.shape[-1]) mask = self._build_mask(action.shape[-1])
new_transition[TransitionKey.ACTION] = to_relative_actions(action, reference_state, mask) new_transition[TransitionKey.ACTION] = to_relative_actions(
action,
reference_state,
mask,
pose_representation=self.pose_representation,
se3_pose_groups=self.se3_pose_groups,
)
return new_transition return new_transition
def get_cached_state(self) -> torch.Tensor | None: def get_cached_state(self) -> torch.Tensor | None:
@@ -155,6 +312,8 @@ class RelativeActionsProcessorStep(ProcessorStep):
"enabled": self.enabled, "enabled": self.enabled,
"exclude_joints": self.exclude_joints, "exclude_joints": self.exclude_joints,
"action_names": self.action_names, "action_names": self.action_names,
"pose_representation": self.pose_representation,
"se3_pose_groups": self.se3_pose_groups,
} }
def transform_features( def transform_features(
@@ -203,7 +362,13 @@ class AbsoluteActionsProcessorStep(ProcessorStep):
return new_transition return new_transition
mask = self.relative_step._build_mask(action.shape[-1]) mask = self.relative_step._build_mask(action.shape[-1])
new_transition[TransitionKey.ACTION] = to_absolute_actions(action, cached_state, mask) new_transition[TransitionKey.ACTION] = to_absolute_actions(
action,
cached_state,
mask,
pose_representation=self.relative_step.pose_representation,
se3_pose_groups=self.relative_step.se3_pose_groups,
)
return new_transition return new_transition
def get_config(self) -> dict[str, Any]: def get_config(self) -> dict[str, Any]:
@@ -325,6 +325,8 @@ class RecomputeStatsConfig(OperationConfig):
relative_exclude_joints: list[str] | None = None relative_exclude_joints: list[str] | None = None
chunk_size: int = 50 chunk_size: int = 50
num_workers: int = 0 num_workers: int = 0
relative_pose_representation: str = "componentwise"
relative_se3_pose_groups: list[list[int]] | None = None
overwrite: bool = False overwrite: bool = False
@@ -698,6 +700,8 @@ def handle_recompute_stats(cfg: EditDatasetConfig) -> None:
relative_exclude_joints=cfg.operation.relative_exclude_joints, relative_exclude_joints=cfg.operation.relative_exclude_joints,
chunk_size=cfg.operation.chunk_size, chunk_size=cfg.operation.chunk_size,
num_workers=cfg.operation.num_workers, num_workers=cfg.operation.num_workers,
relative_pose_representation=cfg.operation.relative_pose_representation,
relative_se3_pose_groups=cfg.operation.relative_se3_pose_groups,
) )
logging.info(f"Stats written to {dataset.root}") logging.info(f"Stats written to {dataset.root}")
@@ -14,6 +14,8 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from math import pi
import numpy as np import numpy as np
import pytest import pytest
import torch import torch
@@ -114,6 +116,26 @@ def test_state_history_can_be_relative_with_absolute_gripper():
) )
def test_state_history_can_use_se3_composition_with_absolute_gripper():
state_history = torch.tensor(
[[[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.2], [0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.4]]]
)
step = Pi05FlattenStateHistoryProcessorStep(
history_steps=2,
max_state_dim=14,
relative=True,
exclude_joints=["gripper"],
state_names=["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
pose_representation="se3",
se3_pose_groups=[list(range(6))],
)
result = step(_transition(torch.zeros(1, 2, 7), state_history))
expected = torch.tensor([[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.2, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.4]])
torch.testing.assert_close(result[TransitionKey.OBSERVATION][OBS_STATE], expected, atol=1e-6, rtol=1e-6)
def test_inference_state_history_is_rolled_and_reset(): def test_inference_state_history_is_rolled_and_reset():
step = Pi05StateFromActionProcessorStep(enabled=True, history_steps=2) step = Pi05StateFromActionProcessorStep(enabled=True, history_steps=2)
@@ -173,3 +195,33 @@ def test_relative_state_history_stats_match_processor_representation():
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]]) 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)) np.testing.assert_allclose(stats["mean"], expected.mean(axis=0))
def test_se3_relative_action_stats_use_reference_frame():
actions = np.asarray(
[
[0.0, 0.0, 0.0, 0.0, 0.0, pi / 2, 0.2],
[0.0, 1.0, 0.0, 0.0, 0.0, pi / 2, 0.3],
],
dtype=np.float32,
)
dataset = {"action": actions, "episode_index": np.zeros(2, dtype=np.int64)}
features = {
"action": {
"shape": [7],
"names": ["x", "y", "z", "rx", "ry", "rz", "gripper_width"],
}
}
stats = compute_relative_action_stats(
dataset,
features,
chunk_size=2,
exclude_joints=["gripper"],
state_from_action=True,
pose_representation="se3",
se3_pose_groups=[list(range(6))],
)
np.testing.assert_allclose(stats["mean"][:3], [0.5, 0.0, 0.0], atol=1e-6)
np.testing.assert_allclose(stats["mean"][6], 0.25, atol=1e-6)
@@ -0,0 +1,70 @@
import math
import pytest
import torch
from lerobot.processor.relative_action_processor import (
to_absolute_actions,
to_absolute_se3_pose,
to_relative_actions,
to_relative_se3_pose,
)
POSE_GROUP = [list(range(6))]
def test_se3_translation_is_expressed_in_reference_frame():
reference = torch.tensor([[1.0, 2.0, 3.0, 0.0, 0.0, math.pi / 2]])
target = torch.tensor([[1.0, 3.0, 3.0, 0.0, 0.0, math.pi / 2]])
relative = to_relative_se3_pose(target, reference)
torch.testing.assert_close(relative, torch.tensor([[1.0, 0.0, 0.0, 0.0, 0.0, 0.0]]), atol=1e-6, rtol=1e-6)
def test_se3_pose_roundtrip_for_batched_chunks():
torch.manual_seed(0)
reference = torch.randn(4, 6)
reference[:, 3:] *= 0.8
target = torch.randn(4, 11, 6)
target[..., 3:] *= 0.8
relative = to_relative_se3_pose(target, reference.unsqueeze(1))
recovered = to_absolute_se3_pose(relative, reference.unsqueeze(1))
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
def test_mixed_se3_pose_and_absolute_gripper_roundtrip():
reference = torch.tensor([[0.2, -0.1, 0.4, 0.1, 0.2, -0.3, 0.06]])
target = torch.tensor([[[0.3, 0.2, 0.5, -0.2, 0.1, 0.4, 0.03], [0.1, -0.3, 0.2, 0.5, -0.1, 0.2, 0.05]]])
mask = [True, True, True, True, True, True, False]
relative = to_relative_actions(
target,
reference,
mask,
pose_representation="se3",
se3_pose_groups=POSE_GROUP,
)
recovered = to_absolute_actions(
relative,
reference,
mask,
pose_representation="se3",
se3_pose_groups=POSE_GROUP,
)
torch.testing.assert_close(relative[..., 6], target[..., 6])
torch.testing.assert_close(recovered, target, atol=2e-5, rtol=2e-5)
def test_se3_pose_group_cannot_be_partially_relative():
with pytest.raises(ValueError, match="wholly relative or wholly absolute"):
to_relative_actions(
torch.zeros(1, 7),
torch.zeros(1, 7),
[True, True, True, False, False, False, False],
pose_representation="se3",
se3_pose_groups=POSE_GROUP,
)