mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 18:26:11 +00:00
feat(pi05): compose relative poses in SE3
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user