se3labs-interface 0.0.1a1__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- se3labs_interface-0.0.1a1/PKG-INFO +14 -0
- se3labs_interface-0.0.1a1/pyproject.toml +37 -0
- se3labs_interface-0.0.1a1/setup.cfg +4 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/__init__.py +0 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/eval/__init__.py +23 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/eval/arm.py +169 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/eval/camera.py +65 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/eval/job.py +22 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/eval/policy.py +32 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/eval/repository.py +27 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/eval/station.py +179 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/eval/task.py +177 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/util/__init__.py +0 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/util/validation.py +16 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/__init__.py +21 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/codec.py +267 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/job_status.proto +28 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/job_status_pb2.py +40 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/job_status_pb2.pyi +29 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/job_status_pb2_grpc.py +102 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/policy_session.proto +243 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/policy_session_pb2.py +120 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/policy_session_pb2.pyi +353 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/policy_session_pb2_grpc.py +117 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/station_link.proto +49 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/station_link_pb2.py +48 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/station_link_pb2.pyi +49 -0
- se3labs_interface-0.0.1a1/src/se3labs/interface/wire/station_link_pb2_grpc.py +107 -0
- se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/PKG-INFO +14 -0
- se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/SOURCES.txt +33 -0
- se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/dependency_links.txt +1 -0
- se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/requires.txt +10 -0
- se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/top_level.txt +1 -0
- se3labs_interface-0.0.1a1/tests/test_spec.py +239 -0
- se3labs_interface-0.0.1a1/tests/test_wire.py +148 -0
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: se3labs-interface
|
|
3
|
+
Version: 0.0.1a1
|
|
4
|
+
Summary: SE3 Labs SDK shared interface
|
|
5
|
+
Requires-Python: >=3.10
|
|
6
|
+
Requires-Dist: numpy>=2.2
|
|
7
|
+
Requires-Dist: jaxtyping>=0.3
|
|
8
|
+
Requires-Dist: beartype>=0.22
|
|
9
|
+
Requires-Dist: grpcio>=1.83.1
|
|
10
|
+
Requires-Dist: protobuf>=7.35.1
|
|
11
|
+
Requires-Dist: pillow>=10
|
|
12
|
+
Provides-Extra: dev
|
|
13
|
+
Requires-Dist: pytest; extra == "dev"
|
|
14
|
+
Requires-Dist: pyright; extra == "dev"
|
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
requires = ["setuptools>=61"]
|
|
3
|
+
build-backend = "setuptools.build_meta"
|
|
4
|
+
|
|
5
|
+
[project]
|
|
6
|
+
name = "se3labs-interface"
|
|
7
|
+
version = "0.0.1a1"
|
|
8
|
+
description = "SE3 Labs SDK shared interface"
|
|
9
|
+
requires-python = ">=3.10"
|
|
10
|
+
dependencies = [
|
|
11
|
+
"numpy>=2.2",
|
|
12
|
+
"jaxtyping>=0.3",
|
|
13
|
+
"beartype>=0.22",
|
|
14
|
+
"grpcio>=1.83.1",
|
|
15
|
+
"protobuf>=7.35.1",
|
|
16
|
+
"pillow>=10",
|
|
17
|
+
]
|
|
18
|
+
|
|
19
|
+
[project.optional-dependencies]
|
|
20
|
+
dev = [
|
|
21
|
+
"pytest",
|
|
22
|
+
"pyright",
|
|
23
|
+
]
|
|
24
|
+
|
|
25
|
+
[tool.setuptools.packages.find]
|
|
26
|
+
where = ["src"]
|
|
27
|
+
|
|
28
|
+
[tool.setuptools.package-data]
|
|
29
|
+
"se3labs.interface.wire" = ["*.proto", "*.pyi"]
|
|
30
|
+
|
|
31
|
+
[tool.pytest.ini_options]
|
|
32
|
+
testpaths = ["tests"]
|
|
33
|
+
|
|
34
|
+
[tool.pyright]
|
|
35
|
+
include = ["src", "tests"]
|
|
36
|
+
exclude = ["**/__pycache__", "**/*_pb2.py", "**/*_pb2_grpc.py"]
|
|
37
|
+
typeCheckingMode = "standard"
|
|
File without changes
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
from . import arm, camera, job, policy, station, task
|
|
2
|
+
from .arm import (
|
|
3
|
+
MotorizedArm_ArmType,
|
|
4
|
+
MotorizedArm_GripperType,
|
|
5
|
+
MotorizedArmSpec,
|
|
6
|
+
MotorizedArm_Observation,
|
|
7
|
+
MotorizedArm_JointPose,
|
|
8
|
+
MotorizedArm_ActionChunk,
|
|
9
|
+
)
|
|
10
|
+
from .camera import Camera_DeviceType, Camera_DeviceSpec, Camera_Observation
|
|
11
|
+
from .station import (
|
|
12
|
+
Station_BaseSpec,
|
|
13
|
+
Station_Spec,
|
|
14
|
+
Bimanual_Station_Spec,
|
|
15
|
+
Single_Arm_Station_Spec,
|
|
16
|
+
Station_Observation,
|
|
17
|
+
Station_JointPose,
|
|
18
|
+
Station_ActionChunk,
|
|
19
|
+
)
|
|
20
|
+
from .task import Task_ScoreKind, Task_ScoreChoice, Task_ScoreItem, Task_ScoreAnswers, Task_ScoreSheet, Task_Spec
|
|
21
|
+
from .policy import Policy
|
|
22
|
+
from .job import JobRequest
|
|
23
|
+
from .repository import Bimanual_YAM_Station
|
|
@@ -0,0 +1,169 @@
|
|
|
1
|
+
"""Motorized arm hardware spec and the per-arm observation / action containers.
|
|
2
|
+
|
|
3
|
+
Conventions (all arrays are ``np.float32``):
|
|
4
|
+
|
|
5
|
+
* joint angles in radians, joint velocities in rad/s, joint torques in N*m,
|
|
6
|
+
joint temperatures in degrees Celsius.
|
|
7
|
+
* ``gripper_pos`` is a normalized open fraction in ``[0, 1]`` where ``0`` is
|
|
8
|
+
fully closed and ``1`` is fully open. The station maps this to the physical
|
|
9
|
+
gripper travel, so policies never see device units.
|
|
10
|
+
* Action targets are **absolute** joint positions.
|
|
11
|
+
"""
|
|
12
|
+
import numpy as np
|
|
13
|
+
import jaxtyping as jt
|
|
14
|
+
from beartype import beartype
|
|
15
|
+
from enum import Enum
|
|
16
|
+
from dataclasses import dataclass
|
|
17
|
+
|
|
18
|
+
from ..util.validation import check_hint
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class MotorizedArm_ArmType(Enum):
|
|
22
|
+
YAM_6DoF = "yam_6dof"
|
|
23
|
+
"""A YAM 6-DoF arm.
|
|
24
|
+
|
|
25
|
+
Reference: https://i2rt.com/products/yam-6-dof-arm
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
@property
|
|
29
|
+
def n_dof(self) -> int:
|
|
30
|
+
return _ARM_N_DOF[self]
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
_ARM_N_DOF: dict[MotorizedArm_ArmType, int] = {
|
|
34
|
+
MotorizedArm_ArmType.YAM_6DoF: 6,
|
|
35
|
+
}
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class MotorizedArm_GripperType(Enum):
|
|
39
|
+
I2RT_linear_gripper_1DoF = "i2rt_linear_gripper_1dof"
|
|
40
|
+
"""I2RT linear gripper with 1 DoF.
|
|
41
|
+
|
|
42
|
+
Reference: https://i2rt.com/products/yam-gripper-1
|
|
43
|
+
"""
|
|
44
|
+
|
|
45
|
+
@property
|
|
46
|
+
def n_dof(self) -> int:
|
|
47
|
+
return _GRIPPER_N_DOF[self]
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
_GRIPPER_N_DOF: dict[MotorizedArm_GripperType, int] = {
|
|
51
|
+
MotorizedArm_GripperType.I2RT_linear_gripper_1DoF: 1,
|
|
52
|
+
}
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
@dataclass(frozen=True)
|
|
56
|
+
class MotorizedArmSpec:
|
|
57
|
+
arm : MotorizedArm_ArmType
|
|
58
|
+
gripper : MotorizedArm_GripperType
|
|
59
|
+
|
|
60
|
+
refresh_rate: int
|
|
61
|
+
"""Rate at which observations are published, unit: [Hz]."""
|
|
62
|
+
|
|
63
|
+
control_rate: int
|
|
64
|
+
"""Rate at which joint position targets are sent to the motors, unit: [Hz].
|
|
65
|
+
|
|
66
|
+
Action chunks emitted at a lower rate are interpolated up to this rate by the station.
|
|
67
|
+
"""
|
|
68
|
+
|
|
69
|
+
@property
|
|
70
|
+
def n_arm_dof(self) -> int:
|
|
71
|
+
return self.arm.n_dof
|
|
72
|
+
|
|
73
|
+
@property
|
|
74
|
+
def n_gripper_dof(self) -> int:
|
|
75
|
+
return self.gripper.n_dof
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
@jt.jaxtyped(typechecker=beartype)
|
|
79
|
+
@dataclass(frozen=True)
|
|
80
|
+
class MotorizedArm_Observation:
|
|
81
|
+
"""One observation tick of a single arm. Every channel is always populated.
|
|
82
|
+
|
|
83
|
+
Construction enforces dtype / rank and that every joint channel shares ``n_arm_dof``;
|
|
84
|
+
``validate`` binds ``n_arm_dof`` / ``n_gripper_dof`` to a concrete ``MotorizedArmSpec``.
|
|
85
|
+
"""
|
|
86
|
+
|
|
87
|
+
timestamp : float
|
|
88
|
+
"""Monotonic capture time on the station clock, unit: [s]."""
|
|
89
|
+
|
|
90
|
+
joint_angle_q : jt.Float32[np.ndarray, "n_arm_dof"]
|
|
91
|
+
"""Joint positions, unit: [rad]."""
|
|
92
|
+
joint_angle_vel_dq: jt.Float32[np.ndarray, "n_arm_dof"]
|
|
93
|
+
"""Joint velocities, unit: [rad/s]."""
|
|
94
|
+
joint_torque : jt.Float32[np.ndarray, "n_arm_dof"]
|
|
95
|
+
"""Joint torques, unit: [N*m]."""
|
|
96
|
+
joint_temperature : jt.Float32[np.ndarray, "n_arm_dof"]
|
|
97
|
+
"""Motor temperatures, unit: [degC]."""
|
|
98
|
+
gripper_pos : jt.Float32[np.ndarray, "n_gripper_dof"]
|
|
99
|
+
"""Normalized gripper open fraction in [0, 1]; 0 = closed, 1 = open."""
|
|
100
|
+
|
|
101
|
+
def validate(self, spec: MotorizedArmSpec) -> None:
|
|
102
|
+
joint = jt.Float32[np.ndarray, f"{spec.n_arm_dof}"]
|
|
103
|
+
check_hint("joint_angle_q" , self.joint_angle_q , joint)
|
|
104
|
+
check_hint("joint_angle_vel_dq", self.joint_angle_vel_dq, joint)
|
|
105
|
+
check_hint("joint_torque" , self.joint_torque , joint)
|
|
106
|
+
check_hint("joint_temperature" , self.joint_temperature , joint)
|
|
107
|
+
check_hint("gripper_pos" , self.gripper_pos , jt.Float32[np.ndarray, f"{spec.n_gripper_dof}"])
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
@jt.jaxtyped(typechecker=beartype)
|
|
111
|
+
@dataclass(frozen=True)
|
|
112
|
+
class MotorizedArm_JointPose:
|
|
113
|
+
"""One absolute joint configuration of a single arm, e.g. the pose the arm walks to before an episode."""
|
|
114
|
+
|
|
115
|
+
joint_angle_q: jt.Float32[np.ndarray, "n_arm_dof"]
|
|
116
|
+
"""Absolute joint positions, unit: [rad]."""
|
|
117
|
+
|
|
118
|
+
gripper_pos : jt.Float32[np.ndarray, "n_gripper_dof"]
|
|
119
|
+
"""Normalized gripper open fraction in [0, 1]; 0 = closed, 1 = open."""
|
|
120
|
+
|
|
121
|
+
def __post_init__(self) -> None:
|
|
122
|
+
if np.any(self.gripper_pos < 0.0) or np.any(self.gripper_pos > 1.0):
|
|
123
|
+
raise ValueError("gripper_pos must lie in [0, 1]")
|
|
124
|
+
|
|
125
|
+
def validate(self, spec: MotorizedArmSpec) -> None:
|
|
126
|
+
"""Bind the dims to a concrete arm spec."""
|
|
127
|
+
check_hint("joint_angle_q", self.joint_angle_q, jt.Float32[np.ndarray, f"{spec.n_arm_dof}"])
|
|
128
|
+
check_hint("gripper_pos" , self.gripper_pos , jt.Float32[np.ndarray, f"{spec.n_gripper_dof}"])
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
@jt.jaxtyped(typechecker=beartype)
|
|
132
|
+
@dataclass(frozen=True)
|
|
133
|
+
class MotorizedArm_ActionChunk:
|
|
134
|
+
"""A horizon of absolute joint position targets for a single arm.
|
|
135
|
+
|
|
136
|
+
Construction enforces dtype / rank and a shared ``horizon`` between joints and gripper;
|
|
137
|
+
``validate`` binds the dims to a concrete ``MotorizedArmSpec`` and horizon.
|
|
138
|
+
|
|
139
|
+
Row ``i`` is the target ``i / rate`` seconds after the chunk's reference timestamp
|
|
140
|
+
(both carried by the enclosing ``station.Station_ActionChunk``).
|
|
141
|
+
"""
|
|
142
|
+
|
|
143
|
+
joint_angle_q: jt.Float32[np.ndarray, "horizon n_arm_dof"]
|
|
144
|
+
"""Absolute joint position targets, unit: [rad]."""
|
|
145
|
+
|
|
146
|
+
gripper_pos : jt.Float32[np.ndarray, "horizon n_gripper_dof"]
|
|
147
|
+
"""Normalized gripper open fraction in [0, 1]; 0 = closed, 1 = open."""
|
|
148
|
+
|
|
149
|
+
def __post_init__(self) -> None:
|
|
150
|
+
if np.any(self.gripper_pos < 0.0) or np.any(self.gripper_pos > 1.0):
|
|
151
|
+
raise ValueError("gripper_pos must lie in [0, 1]")
|
|
152
|
+
|
|
153
|
+
@property
|
|
154
|
+
def horizon(self) -> int:
|
|
155
|
+
return int(self.joint_angle_q.shape[0])
|
|
156
|
+
|
|
157
|
+
def __len__(self) -> int:
|
|
158
|
+
return self.horizon
|
|
159
|
+
|
|
160
|
+
def __getitem__(self, index: int) -> MotorizedArm_JointPose:
|
|
161
|
+
"""Row ``index`` of the chunk as a single joint pose. Negative indices count from the end."""
|
|
162
|
+
if not -self.horizon <= index < self.horizon:
|
|
163
|
+
raise IndexError(f"chunk index {index} out of range for horizon {self.horizon}")
|
|
164
|
+
return MotorizedArm_JointPose(self.joint_angle_q[index], self.gripper_pos[index])
|
|
165
|
+
|
|
166
|
+
def validate(self, spec: MotorizedArmSpec, horizon: int) -> None:
|
|
167
|
+
"""Bind the dims to a concrete arm spec and horizon."""
|
|
168
|
+
check_hint("joint_angle_q", self.joint_angle_q, jt.Float32[np.ndarray, f"{horizon} {spec.n_arm_dof}"])
|
|
169
|
+
check_hint("gripper_pos" , self.gripper_pos , jt.Float32[np.ndarray, f"{horizon} {spec.n_gripper_dof}"])
|
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import jaxtyping as jt
|
|
3
|
+
from enum import Enum
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from beartype import beartype
|
|
6
|
+
|
|
7
|
+
from ..util.validation import check_hint
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class Camera_DeviceType(Enum):
|
|
11
|
+
Realsense_D405 = "realsense_d405"
|
|
12
|
+
"""Realsense D405 Stereo Camera. Both imagers are RGB, so one observation carries 2 views (left, right)."""
|
|
13
|
+
|
|
14
|
+
@property
|
|
15
|
+
def n_views(self) -> int:
|
|
16
|
+
"""Number of frames delivered per observation (1 for mono, 2 for a stereo pair)."""
|
|
17
|
+
return _CAMERA_N_VIEWS[self]
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
_CAMERA_N_VIEWS: dict[Camera_DeviceType, int] = {
|
|
21
|
+
Camera_DeviceType.Realsense_D405: 2,
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@dataclass(frozen=True)
|
|
26
|
+
class Camera_DeviceSpec:
|
|
27
|
+
device_type : Camera_DeviceType
|
|
28
|
+
"""Device type of the camera"""
|
|
29
|
+
|
|
30
|
+
resolution : tuple[int, int]
|
|
31
|
+
"""Resolution of each view of the camera stream, in tuple of (height, width), unit: [px]"""
|
|
32
|
+
|
|
33
|
+
refresh_rate : int
|
|
34
|
+
"""Refresh rate of the camera, unit: [fps]"""
|
|
35
|
+
|
|
36
|
+
jpeg_compression: int | None
|
|
37
|
+
"""jpeg compression quality value if specified, if None then no compression will be made."""
|
|
38
|
+
|
|
39
|
+
@property
|
|
40
|
+
def n_views(self) -> int:
|
|
41
|
+
return self.device_type.n_views
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
@jt.jaxtyped(typechecker=beartype)
|
|
45
|
+
@dataclass(frozen=True)
|
|
46
|
+
class Camera_Observation:
|
|
47
|
+
"""One decoded observation of a camera: ``b`` RGB views captured together (``b = 2`` for a stereo pair).
|
|
48
|
+
|
|
49
|
+
Frames are delivered already decoded regardless of ``jpeg_compression``. Construction enforces
|
|
50
|
+
``uint8`` BHWC with 3 channels; ``validate`` binds (b, h, w) to the camera spec.
|
|
51
|
+
"""
|
|
52
|
+
|
|
53
|
+
timestamp: float
|
|
54
|
+
"""Monotonic capture time on the station clock, unit: [s]. Shared by all views."""
|
|
55
|
+
|
|
56
|
+
rgb : jt.UInt8[np.ndarray, "b h w 3"]
|
|
57
|
+
"""Views in device order (stereo: index 0 = left, 1 = right)."""
|
|
58
|
+
|
|
59
|
+
@property
|
|
60
|
+
def n_views(self) -> int:
|
|
61
|
+
return int(self.rgb.shape[0])
|
|
62
|
+
|
|
63
|
+
def validate(self, spec: Camera_DeviceSpec) -> None:
|
|
64
|
+
h, w = spec.resolution
|
|
65
|
+
check_hint("rgb", self.rgb, jt.UInt8[np.ndarray, f"{spec.n_views} {h} {w} 3"])
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
"""What a customer submits to a station: the task, the policy, and how many
|
|
2
|
+
episodes to run. The station executes the request as uploaded; it does not
|
|
3
|
+
get to say what is being evaluated."""
|
|
4
|
+
|
|
5
|
+
from __future__ import annotations
|
|
6
|
+
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
|
|
9
|
+
from .task import Task_Spec
|
|
10
|
+
|
|
11
|
+
from .policy import Policy
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@dataclass(frozen=True)
|
|
15
|
+
class JobRequest:
|
|
16
|
+
task : Task_Spec
|
|
17
|
+
policy : Policy
|
|
18
|
+
episodes: int
|
|
19
|
+
|
|
20
|
+
def __post_init__(self) -> None:
|
|
21
|
+
if self.episodes <= 0:
|
|
22
|
+
raise ValueError("a job needs at least one episode")
|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
"""The policy contract: what a caller hands to ``evaluate(station, task, policy)``.
|
|
2
|
+
|
|
3
|
+
Lifecycle, driven by the station, once per episode:
|
|
4
|
+
|
|
5
|
+
1. ``initialize(station, task)``: return the joint pose the arms walk to before the episode starts.
|
|
6
|
+
Load weights and read the task instruction here.
|
|
7
|
+
2. ``infer(observation)``, repeated on every observation tick until the episode ends. Return a new
|
|
8
|
+
chunk, or ``None`` to keep executing the current one. A new chunk replaces the remainder of the
|
|
9
|
+
current one.
|
|
10
|
+
3. ``reset()`` after the episode ends. Clear observation history and any unexecuted chunk so the
|
|
11
|
+
next ``initialize`` starts clean.
|
|
12
|
+
|
|
13
|
+
Anything about how the policy runs (its rate, chunk length, replanning, image preprocessing) lives
|
|
14
|
+
inside these methods and is not declared to the SDK.
|
|
15
|
+
"""
|
|
16
|
+
from abc import ABC, abstractmethod
|
|
17
|
+
|
|
18
|
+
from .station import Station_Spec, Station_Observation, Station_ActionChunk, Station_JointPose
|
|
19
|
+
from .task import Task_Spec
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class Policy(ABC):
|
|
23
|
+
@abstractmethod
|
|
24
|
+
def initialize(self, station: Station_Spec, task: Task_Spec) -> Station_JointPose:
|
|
25
|
+
"""Prepare for one episode of ``task`` on ``station`` and return the pose the arms start from."""
|
|
26
|
+
|
|
27
|
+
@abstractmethod
|
|
28
|
+
def infer(self, observation: Station_Observation) -> Station_ActionChunk | None:
|
|
29
|
+
"""A new action chunk for this tick, or ``None`` to continue executing the current chunk."""
|
|
30
|
+
|
|
31
|
+
def reset(self) -> None:
|
|
32
|
+
"""Clear per-episode state. Called once after the episode ends. No-op by default."""
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
"""Registry of known station specs."""
|
|
2
|
+
from . import arm, camera
|
|
3
|
+
from .station import Bimanual_Station_Spec, Station_BaseSpec
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
Bimanual_YAM_Station = Bimanual_Station_Spec(
|
|
7
|
+
base=Station_BaseSpec.Static,
|
|
8
|
+
arm =arm.MotorizedArmSpec(
|
|
9
|
+
arm =arm.MotorizedArm_ArmType.YAM_6DoF,
|
|
10
|
+
gripper =arm.MotorizedArm_GripperType.I2RT_linear_gripper_1DoF,
|
|
11
|
+
refresh_rate=100,
|
|
12
|
+
control_rate=100,
|
|
13
|
+
),
|
|
14
|
+
wrist_camera=camera.Camera_DeviceSpec(
|
|
15
|
+
device_type =camera.Camera_DeviceType.Realsense_D405,
|
|
16
|
+
resolution =(360, 640), # (h, w); 2x downsample for 720p.
|
|
17
|
+
refresh_rate =30,
|
|
18
|
+
jpeg_compression=80
|
|
19
|
+
),
|
|
20
|
+
topdown_camera=camera.Camera_DeviceSpec(
|
|
21
|
+
device_type =camera.Camera_DeviceType.Realsense_D405,
|
|
22
|
+
resolution =(360, 640), # (h, w); 2x downsample for 720p.
|
|
23
|
+
refresh_rate =30,
|
|
24
|
+
jpeg_compression=80
|
|
25
|
+
)
|
|
26
|
+
)
|
|
27
|
+
"""Static bimanual YAM station: two YAM 6-DoF arms with I2RT grippers, two wrist D405s, one top-down D405."""
|
|
@@ -0,0 +1,179 @@
|
|
|
1
|
+
"""Station specs (hardware capability) and the station-level observation / action containers.
|
|
2
|
+
|
|
3
|
+
A station is a named collection of arms and cameras. Observations and actions are keyed by
|
|
4
|
+
those names so that a bimanual and a single-arm station share one container type.
|
|
5
|
+
|
|
6
|
+
The policy side of the contract is ``policy.Policy``: it consumes ``Station_Observation`` ticks and
|
|
7
|
+
produces ``Station_ActionChunk``s. Anything about how the policy runs (its rate, chunk length,
|
|
8
|
+
replanning, image preprocessing) is decided inside the policy, not declared to the SDK.
|
|
9
|
+
"""
|
|
10
|
+
from abc import ABC, abstractmethod
|
|
11
|
+
from enum import Enum
|
|
12
|
+
from dataclasses import dataclass
|
|
13
|
+
|
|
14
|
+
from . import arm as arm_types
|
|
15
|
+
from . import camera
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class Station_BaseSpec(Enum):
|
|
19
|
+
Static = "static"
|
|
20
|
+
"""Static manipulation base (the arm is mounted on a fixed, non-mobile base)."""
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class Station_Spec(ABC):
|
|
24
|
+
base: Station_BaseSpec
|
|
25
|
+
|
|
26
|
+
@property
|
|
27
|
+
@abstractmethod
|
|
28
|
+
def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
|
|
29
|
+
"""Arm name -> spec. Names are the keys of ``Station_Observation.arms`` and ``Station_ActionChunk.arms``."""
|
|
30
|
+
|
|
31
|
+
@property
|
|
32
|
+
@abstractmethod
|
|
33
|
+
def cameras(self) -> dict[str, camera.Camera_DeviceSpec]:
|
|
34
|
+
"""Camera name -> spec. Names are the keys of ``Station_Observation.cameras``."""
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
@dataclass(frozen=True)
|
|
38
|
+
class Bimanual_Station_Spec(Station_Spec):
|
|
39
|
+
base : Station_BaseSpec
|
|
40
|
+
|
|
41
|
+
arm : arm_types.MotorizedArmSpec
|
|
42
|
+
"""Spec shared by the ``left`` and ``right`` arms."""
|
|
43
|
+
|
|
44
|
+
wrist_camera : camera.Camera_DeviceSpec
|
|
45
|
+
"""Spec shared by the ``wrist_left`` and ``wrist_right`` cameras."""
|
|
46
|
+
|
|
47
|
+
topdown_camera: camera.Camera_DeviceSpec
|
|
48
|
+
|
|
49
|
+
@property
|
|
50
|
+
def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
|
|
51
|
+
return {"left": self.arm, "right": self.arm}
|
|
52
|
+
|
|
53
|
+
@property
|
|
54
|
+
def cameras(self) -> dict[str, camera.Camera_DeviceSpec]:
|
|
55
|
+
return {
|
|
56
|
+
"wrist_left" : self.wrist_camera,
|
|
57
|
+
"wrist_right": self.wrist_camera,
|
|
58
|
+
"topdown" : self.topdown_camera,
|
|
59
|
+
}
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
@dataclass(frozen=True)
|
|
63
|
+
class Single_Arm_Station_Spec(Station_Spec):
|
|
64
|
+
base : Station_BaseSpec
|
|
65
|
+
arm : arm_types.MotorizedArmSpec
|
|
66
|
+
wrist_camera : camera.Camera_DeviceSpec
|
|
67
|
+
|
|
68
|
+
@property
|
|
69
|
+
def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
|
|
70
|
+
return {"arm": self.arm}
|
|
71
|
+
|
|
72
|
+
@property
|
|
73
|
+
def cameras(self) -> dict[str, camera.Camera_DeviceSpec]:
|
|
74
|
+
return {"wrist": self.wrist_camera}
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
@dataclass(frozen=True)
|
|
78
|
+
class Station_Observation:
|
|
79
|
+
"""One policy-side observation: the latest tick of every arm and camera.
|
|
80
|
+
|
|
81
|
+
Arms and cameras run at different rates, so each entry carries its own timestamp; consumers that
|
|
82
|
+
need temporal alignment (world models, chunked policies) should use those rather than assume sync.
|
|
83
|
+
"""
|
|
84
|
+
|
|
85
|
+
arms : dict[str, arm_types.MotorizedArm_Observation]
|
|
86
|
+
cameras: dict[str, camera.Camera_Observation]
|
|
87
|
+
|
|
88
|
+
def validate(self, station: Station_Spec) -> None:
|
|
89
|
+
"""Check every arm / camera of ``station`` is present with the shapes its spec promises."""
|
|
90
|
+
missing_arms = set(station.arms) - set(self.arms)
|
|
91
|
+
if missing_arms:
|
|
92
|
+
raise ValueError(f"observation missing arms: {sorted(missing_arms)}")
|
|
93
|
+
for name, spec in station.arms.items():
|
|
94
|
+
try:
|
|
95
|
+
self.arms[name].validate(spec)
|
|
96
|
+
except (TypeError, ValueError) as e:
|
|
97
|
+
raise type(e)(f"arm '{name}': {e}") from None
|
|
98
|
+
|
|
99
|
+
missing_cams = set(station.cameras) - set(self.cameras)
|
|
100
|
+
if missing_cams:
|
|
101
|
+
raise ValueError(f"observation missing cameras: {sorted(missing_cams)}")
|
|
102
|
+
for name, spec in station.cameras.items():
|
|
103
|
+
try:
|
|
104
|
+
self.cameras[name].validate(spec)
|
|
105
|
+
except (TypeError, ValueError) as e:
|
|
106
|
+
raise type(e)(f"camera '{name}': {e}") from None
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
@dataclass(frozen=True)
|
|
110
|
+
class Station_JointPose:
|
|
111
|
+
"""One absolute joint configuration for every arm of a station."""
|
|
112
|
+
|
|
113
|
+
arms: dict[str, arm_types.MotorizedArm_JointPose]
|
|
114
|
+
|
|
115
|
+
def __post_init__(self) -> None:
|
|
116
|
+
if not self.arms:
|
|
117
|
+
raise ValueError("joint pose has no arms")
|
|
118
|
+
|
|
119
|
+
def validate(self, station: Station_Spec) -> None:
|
|
120
|
+
"""Check the pose names exactly the station's arms with the right shapes."""
|
|
121
|
+
if set(self.arms) != set(station.arms):
|
|
122
|
+
raise ValueError(f"pose arms {sorted(self.arms)} != station arms {sorted(station.arms)}")
|
|
123
|
+
for name, spec in station.arms.items():
|
|
124
|
+
try:
|
|
125
|
+
self.arms[name].validate(spec)
|
|
126
|
+
except (TypeError, ValueError) as e:
|
|
127
|
+
raise type(e)(f"arm '{name}': {e}") from None
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
@dataclass(frozen=True)
|
|
131
|
+
class Station_ActionChunk:
|
|
132
|
+
"""A chunk of ``horizon`` absolute joint-position targets for every arm of a station.
|
|
133
|
+
|
|
134
|
+
Row ``i`` of every arm is the target at ``timestamp + i / rate``. The station interpolates rows
|
|
135
|
+
up to each arm's ``control_rate`` and executes them until the policy returns a new chunk, which
|
|
136
|
+
replaces the remainder of the current one.
|
|
137
|
+
"""
|
|
138
|
+
|
|
139
|
+
timestamp: float
|
|
140
|
+
"""Station-clock time of the observation this chunk was predicted from, unit: [s]."""
|
|
141
|
+
|
|
142
|
+
rate : float
|
|
143
|
+
"""Spacing between consecutive rows, unit: [Hz]. Chosen by the policy per chunk."""
|
|
144
|
+
|
|
145
|
+
arms : dict[str, arm_types.MotorizedArm_ActionChunk]
|
|
146
|
+
|
|
147
|
+
def __post_init__(self) -> None:
|
|
148
|
+
if self.rate <= 0:
|
|
149
|
+
raise ValueError(f"rate must be positive, got {self.rate}")
|
|
150
|
+
if not self.arms:
|
|
151
|
+
raise ValueError("action chunk has no arms")
|
|
152
|
+
horizons = {a.horizon for a in self.arms.values()}
|
|
153
|
+
if len(horizons) != 1:
|
|
154
|
+
raise ValueError(f"inconsistent horizons across arms: {horizons}")
|
|
155
|
+
|
|
156
|
+
@property
|
|
157
|
+
def horizon(self) -> int:
|
|
158
|
+
return next(iter(self.arms.values())).horizon
|
|
159
|
+
|
|
160
|
+
def __len__(self) -> int:
|
|
161
|
+
return self.horizon
|
|
162
|
+
|
|
163
|
+
def __getitem__(self, index: int) -> Station_JointPose:
|
|
164
|
+
"""Row ``index`` of the chunk, for every arm, as one station pose. Negative indices count from the end."""
|
|
165
|
+
return Station_JointPose({name: chunk[index] for name, chunk in self.arms.items()})
|
|
166
|
+
|
|
167
|
+
def validate(self, station: Station_Spec) -> None:
|
|
168
|
+
"""Check the chunk names exactly the station's arms, with the right shapes and a feasible rate."""
|
|
169
|
+
if set(self.arms) != set(station.arms):
|
|
170
|
+
raise ValueError(f"action arms {sorted(self.arms)} != station arms {sorted(station.arms)}")
|
|
171
|
+
h = self.horizon
|
|
172
|
+
for name, spec in station.arms.items():
|
|
173
|
+
if self.rate > spec.control_rate:
|
|
174
|
+
raise ValueError(f"arm '{name}': chunk rate {self.rate} Hz exceeds control_rate {spec.control_rate} Hz")
|
|
175
|
+
try:
|
|
176
|
+
self.arms[name].validate(spec, h)
|
|
177
|
+
except (TypeError, ValueError) as e:
|
|
178
|
+
raise type(e)(f"arm '{name}': {e}") from None
|
|
179
|
+
|