se3labs-interface 0.0.1__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.1/PKG-INFO +14 -0
- se3labs_interface-0.0.1/pyproject.toml +37 -0
- se3labs_interface-0.0.1/setup.cfg +4 -0
- se3labs_interface-0.0.1/src/se3labs/interface/__init__.py +8 -0
- se3labs_interface-0.0.1/src/se3labs/interface/eval/__init__.py +40 -0
- se3labs_interface-0.0.1/src/se3labs/interface/eval/arm.py +194 -0
- se3labs_interface-0.0.1/src/se3labs/interface/eval/camera.py +70 -0
- se3labs_interface-0.0.1/src/se3labs/interface/eval/job.py +20 -0
- se3labs_interface-0.0.1/src/se3labs/interface/eval/policy.py +71 -0
- se3labs_interface-0.0.1/src/se3labs/interface/eval/repository.py +43 -0
- se3labs_interface-0.0.1/src/se3labs/interface/eval/station.py +173 -0
- se3labs_interface-0.0.1/src/se3labs/interface/eval/task.py +164 -0
- se3labs_interface-0.0.1/src/se3labs/interface/util/__init__.py +0 -0
- se3labs_interface-0.0.1/src/se3labs/interface/util/validation.py +16 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/__init__.py +32 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/codec.py +288 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/job_status.proto +43 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/job_status_pb2.py +41 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/job_status_pb2.pyi +30 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/job_status_pb2_grpc.py +102 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/policy_session.proto +320 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/policy_session_pb2.py +124 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/policy_session_pb2.pyi +393 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/policy_session_pb2_grpc.py +117 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/station_link.proto +49 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/station_link_pb2.py +48 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/station_link_pb2.pyi +49 -0
- se3labs_interface-0.0.1/src/se3labs/interface/wire/station_link_pb2_grpc.py +107 -0
- se3labs_interface-0.0.1/src/se3labs_interface.egg-info/PKG-INFO +14 -0
- se3labs_interface-0.0.1/src/se3labs_interface.egg-info/SOURCES.txt +33 -0
- se3labs_interface-0.0.1/src/se3labs_interface.egg-info/dependency_links.txt +1 -0
- se3labs_interface-0.0.1/src/se3labs_interface.egg-info/requires.txt +10 -0
- se3labs_interface-0.0.1/src/se3labs_interface.egg-info/top_level.txt +1 -0
- se3labs_interface-0.0.1/tests/test_spec.py +224 -0
- se3labs_interface-0.0.1/tests/test_wire.py +146 -0
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: se3labs-interface
|
|
3
|
+
Version: 0.0.1
|
|
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.1"
|
|
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"
|
|
@@ -0,0 +1,8 @@
|
|
|
1
|
+
"""@title: Python Interface
|
|
2
|
+
|
|
3
|
+
Shared robot policy contracts and data types, installed by `se3labs-interface`.
|
|
4
|
+
|
|
5
|
+
The `eval` package defines the policies, requests, tasks, and hardware data
|
|
6
|
+
exchanged during evaluations. The SDK installs this package as a dependency.
|
|
7
|
+
Use [`se3labs.sdk`](sdk.md) to connect to a station and run a policy.
|
|
8
|
+
"""
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
"""@title: Evaluation Types
|
|
2
|
+
|
|
3
|
+
Types used to define a robot policy and submit it for evaluation.
|
|
4
|
+
|
|
5
|
+
`policy` defines the policy lifecycle; `job` combines a policy, task, and
|
|
6
|
+
episode count into a request. `task` defines instructions and scoring.
|
|
7
|
+
`station`, `arm`, and `camera` describe the hardware and its observations
|
|
8
|
+
and actions. `repository` contains predefined station specifications.
|
|
9
|
+
|
|
10
|
+
These types can also be imported directly from `se3labs.interface.eval`:
|
|
11
|
+
|
|
12
|
+
```python
|
|
13
|
+
from se3labs.interface.eval import JobRequest, Policy, Task_Spec
|
|
14
|
+
```
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from . import arm, camera, job, policy, station, task
|
|
18
|
+
from .arm import (
|
|
19
|
+
MotorizedArm_ArmType,
|
|
20
|
+
MotorizedArm_GripperType,
|
|
21
|
+
MotorizedArmSpec,
|
|
22
|
+
MotorizedArm_Observation,
|
|
23
|
+
MotorizedArm_JointPose,
|
|
24
|
+
MotorizedArm_ActionChunk,
|
|
25
|
+
)
|
|
26
|
+
from .camera import Camera_DeviceType, Camera_DeviceSpec, Camera_Observation
|
|
27
|
+
from .station import (
|
|
28
|
+
Station_BaseSpec,
|
|
29
|
+
Station_Spec,
|
|
30
|
+
Bimanual_Station_Spec,
|
|
31
|
+
Teleop_Bimanual_Station_Spec,
|
|
32
|
+
Single_Arm_Station_Spec,
|
|
33
|
+
Station_Observation,
|
|
34
|
+
Station_JointPose,
|
|
35
|
+
Station_ActionChunk,
|
|
36
|
+
)
|
|
37
|
+
from .task import Task_ScoreKind, Task_ScoreChoice, Task_ScoreItem, Task_ScoreAnswers, Task_ScoreSheet, Task_Spec
|
|
38
|
+
from .policy import EndEpisode, Policy
|
|
39
|
+
from .job import JobRequest
|
|
40
|
+
from .repository import Bimanual_YAM_Station, Teleop_Bimanual_YAM_Station, Teleop_Bimanual_OpenYAM_Station
|
|
@@ -0,0 +1,194 @@
|
|
|
1
|
+
"""@title: Arms
|
|
2
|
+
|
|
3
|
+
Motorized arm hardware spec and the per-arm observation / action containers.
|
|
4
|
+
|
|
5
|
+
Conventions (all arrays are ``np.float32``):
|
|
6
|
+
|
|
7
|
+
* joint angles in radians, joint velocities in rad/s, joint torques in N*m,
|
|
8
|
+
joint temperatures in degrees Celsius.
|
|
9
|
+
* ``gripper_pos`` is a normalized open fraction in ``[0, 1]`` where ``0`` is
|
|
10
|
+
fully closed and ``1`` is fully open. The station maps this to the physical
|
|
11
|
+
gripper travel, so policies never see device units.
|
|
12
|
+
* Action targets are **absolute** joint positions.
|
|
13
|
+
"""
|
|
14
|
+
import numpy as np
|
|
15
|
+
import jaxtyping as jt
|
|
16
|
+
from beartype import beartype
|
|
17
|
+
from enum import Enum
|
|
18
|
+
from dataclasses import dataclass
|
|
19
|
+
|
|
20
|
+
from ..util.validation import check_hint
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class MotorizedArm_ArmType(Enum):
|
|
24
|
+
YAM_6DoF = "yam_6dof"
|
|
25
|
+
"""A YAM 6-DoF arm.
|
|
26
|
+
|
|
27
|
+
Reference: https://i2rt.com/products/yam-6-dof-arm
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
@property
|
|
31
|
+
def n_dof(self) -> int:
|
|
32
|
+
return _ARM_N_DOF[self]
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
_ARM_N_DOF: dict[MotorizedArm_ArmType, int] = {
|
|
36
|
+
MotorizedArm_ArmType.YAM_6DoF: 6,
|
|
37
|
+
}
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class MotorizedArm_GripperType(Enum):
|
|
41
|
+
I2RT_linear_gripper_1DoF = "i2rt_linear_gripper_1dof"
|
|
42
|
+
"""I2RT linear gripper with 1 DoF.
|
|
43
|
+
|
|
44
|
+
Reference: https://i2rt.com/products/yam-gripper-1
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
I2RT_teaching_handle = "i2rt_teaching_handle"
|
|
48
|
+
"""Passive trigger reporting a normalized opening; no gripper actuator."""
|
|
49
|
+
|
|
50
|
+
@property
|
|
51
|
+
def n_dof(self) -> int:
|
|
52
|
+
return _GRIPPER_N_DOF[self]
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
_GRIPPER_N_DOF: dict[MotorizedArm_GripperType, int] = {
|
|
56
|
+
MotorizedArm_GripperType.I2RT_linear_gripper_1DoF: 1,
|
|
57
|
+
MotorizedArm_GripperType.I2RT_teaching_handle: 1,
|
|
58
|
+
}
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
@dataclass(frozen=True)
|
|
62
|
+
class MotorizedArmSpec:
|
|
63
|
+
arm : MotorizedArm_ArmType
|
|
64
|
+
gripper : MotorizedArm_GripperType
|
|
65
|
+
|
|
66
|
+
refresh_rate: int
|
|
67
|
+
"""Rate at which observations are published, unit: [Hz]."""
|
|
68
|
+
|
|
69
|
+
control_rate: int
|
|
70
|
+
"""Rate at which joint position targets are sent to the motors, unit: [Hz].
|
|
71
|
+
|
|
72
|
+
Action chunks emitted at a lower rate are interpolated up to this rate by the station.
|
|
73
|
+
"""
|
|
74
|
+
|
|
75
|
+
@property
|
|
76
|
+
def n_arm_dof(self) -> int:
|
|
77
|
+
return self.arm.n_dof
|
|
78
|
+
|
|
79
|
+
@property
|
|
80
|
+
def n_gripper_dof(self) -> int:
|
|
81
|
+
return self.gripper.n_dof
|
|
82
|
+
|
|
83
|
+
@property
|
|
84
|
+
def gripper_actuated(self) -> bool:
|
|
85
|
+
return self.gripper != MotorizedArm_GripperType.I2RT_teaching_handle
|
|
86
|
+
|
|
87
|
+
@property
|
|
88
|
+
def actuated_width(self) -> int:
|
|
89
|
+
return self.n_arm_dof + (self.n_gripper_dof if self.gripper_actuated else 0)
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
@jt.jaxtyped(typechecker=beartype)
|
|
93
|
+
@dataclass(frozen=True)
|
|
94
|
+
class MotorizedArm_Observation:
|
|
95
|
+
"""One observation tick of a single arm. Unmeasured channels are NaN.
|
|
96
|
+
|
|
97
|
+
A teaching handle supplies opening only; its gripper velocity, torque,
|
|
98
|
+
and temperature are unavailable.
|
|
99
|
+
|
|
100
|
+
Construction enforces dtype / rank and that every joint channel shares ``n_arm_dof``;
|
|
101
|
+
``validate`` binds ``n_arm_dof`` / ``n_gripper_dof`` to a concrete ``MotorizedArmSpec``.
|
|
102
|
+
"""
|
|
103
|
+
|
|
104
|
+
timestamp : float
|
|
105
|
+
"""Monotonic capture time on the station clock, unit: [s]."""
|
|
106
|
+
|
|
107
|
+
joint_angle_q : jt.Float32[np.ndarray, "n_arm_dof"]
|
|
108
|
+
"""Joint positions, unit: [rad]."""
|
|
109
|
+
joint_angle_vel_dq: jt.Float32[np.ndarray, "n_arm_dof"]
|
|
110
|
+
"""Joint velocities, unit: [rad/s]."""
|
|
111
|
+
joint_torque : jt.Float32[np.ndarray, "n_arm_dof"]
|
|
112
|
+
"""Joint torques, unit: [N*m]."""
|
|
113
|
+
joint_temperature : jt.Float32[np.ndarray, "n_arm_dof"]
|
|
114
|
+
"""Motor temperatures, unit: [degC]."""
|
|
115
|
+
gripper_pos : jt.Float32[np.ndarray, "n_gripper_dof"]
|
|
116
|
+
"""Normalized gripper open fraction in [0, 1]; 0 = closed, 1 = open."""
|
|
117
|
+
gripper_vel : jt.Float32[np.ndarray, "n_gripper_dof"]
|
|
118
|
+
"""Gripper opening velocity, unit: [1/s] (open fraction per second)."""
|
|
119
|
+
gripper_torque : jt.Float32[np.ndarray, "n_gripper_dof"]
|
|
120
|
+
"""Gripper motor torque, unit: [N*m]."""
|
|
121
|
+
gripper_temperature: jt.Float32[np.ndarray, "n_gripper_dof"]
|
|
122
|
+
"""Gripper motor temperature, unit: [degC]."""
|
|
123
|
+
|
|
124
|
+
def validate(self, spec: MotorizedArmSpec) -> None:
|
|
125
|
+
joint = jt.Float32[np.ndarray, f"{spec.n_arm_dof}"]
|
|
126
|
+
check_hint("joint_angle_q" , self.joint_angle_q , joint)
|
|
127
|
+
check_hint("joint_angle_vel_dq", self.joint_angle_vel_dq, joint)
|
|
128
|
+
check_hint("joint_torque" , self.joint_torque , joint)
|
|
129
|
+
check_hint("joint_temperature" , self.joint_temperature , joint)
|
|
130
|
+
gripper = jt.Float32[np.ndarray, f"{spec.n_gripper_dof}"]
|
|
131
|
+
check_hint("gripper_pos" , self.gripper_pos , gripper)
|
|
132
|
+
check_hint("gripper_vel" , self.gripper_vel , gripper)
|
|
133
|
+
check_hint("gripper_torque" , self.gripper_torque , gripper)
|
|
134
|
+
check_hint("gripper_temperature", self.gripper_temperature, gripper)
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
@jt.jaxtyped(typechecker=beartype)
|
|
138
|
+
@dataclass(frozen=True)
|
|
139
|
+
class MotorizedArm_JointPose:
|
|
140
|
+
"""One absolute joint configuration of a single arm, e.g. the pose the arm walks to before an episode."""
|
|
141
|
+
|
|
142
|
+
joint_angle_q: jt.Float32[np.ndarray, "n_arm_dof"]
|
|
143
|
+
"""Absolute joint positions, unit: [rad]."""
|
|
144
|
+
|
|
145
|
+
gripper_pos : jt.Float32[np.ndarray, "n_gripper_dof"]
|
|
146
|
+
"""Normalized gripper open fraction in [0, 1]; 0 = closed, 1 = open."""
|
|
147
|
+
|
|
148
|
+
def __post_init__(self) -> None:
|
|
149
|
+
if np.any(self.gripper_pos < 0.0) or np.any(self.gripper_pos > 1.0):
|
|
150
|
+
raise ValueError("gripper_pos must lie in [0, 1]")
|
|
151
|
+
|
|
152
|
+
def validate(self, spec: MotorizedArmSpec) -> None:
|
|
153
|
+
"""Bind the dims to a concrete arm spec."""
|
|
154
|
+
check_hint("joint_angle_q", self.joint_angle_q, jt.Float32[np.ndarray, f"{spec.n_arm_dof}"])
|
|
155
|
+
check_hint("gripper_pos" , self.gripper_pos , jt.Float32[np.ndarray, f"{spec.n_gripper_dof}"])
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
@jt.jaxtyped(typechecker=beartype)
|
|
159
|
+
@dataclass(frozen=True)
|
|
160
|
+
class MotorizedArm_ActionChunk:
|
|
161
|
+
"""A horizon of absolute joint position targets for a single arm.
|
|
162
|
+
|
|
163
|
+
Construction enforces dtype / rank and a shared ``horizon`` between joints and gripper;
|
|
164
|
+
``validate`` binds the dims to a concrete ``MotorizedArmSpec`` and horizon.
|
|
165
|
+
|
|
166
|
+
Row ``i`` is the target ``i / rate`` seconds after the chunk's reference timestamp
|
|
167
|
+
(both carried by the enclosing ``station.Station_ActionChunk``).
|
|
168
|
+
"""
|
|
169
|
+
|
|
170
|
+
joint_angle_q: jt.Float32[np.ndarray, "horizon n_arm_dof"]
|
|
171
|
+
"""Absolute joint position targets, unit: [rad]."""
|
|
172
|
+
|
|
173
|
+
gripper_pos : jt.Float32[np.ndarray, "horizon n_gripper_dof"]
|
|
174
|
+
"""Normalized gripper open fraction in [0, 1]; 0 = closed, 1 = open."""
|
|
175
|
+
|
|
176
|
+
def __post_init__(self) -> None:
|
|
177
|
+
if np.any(self.gripper_pos < 0.0) or np.any(self.gripper_pos > 1.0):
|
|
178
|
+
raise ValueError("gripper_pos must lie in [0, 1]")
|
|
179
|
+
|
|
180
|
+
@property
|
|
181
|
+
def horizon(self) -> int:
|
|
182
|
+
return int(self.joint_angle_q.shape[0])
|
|
183
|
+
|
|
184
|
+
def __len__(self) -> int:
|
|
185
|
+
return self.horizon
|
|
186
|
+
|
|
187
|
+
def __getitem__(self, index: int) -> MotorizedArm_JointPose:
|
|
188
|
+
"""Row ``index`` of the chunk as a single joint pose. Negative indices count from the end."""
|
|
189
|
+
return MotorizedArm_JointPose(self.joint_angle_q[index], self.gripper_pos[index])
|
|
190
|
+
|
|
191
|
+
def validate(self, spec: MotorizedArmSpec, horizon: int) -> None:
|
|
192
|
+
"""Bind the dims to a concrete arm spec and horizon."""
|
|
193
|
+
check_hint("joint_angle_q", self.joint_angle_q, jt.Float32[np.ndarray, f"{horizon} {spec.n_arm_dof}"])
|
|
194
|
+
check_hint("gripper_pos" , self.gripper_pos , jt.Float32[np.ndarray, f"{horizon} {spec.n_gripper_dof}"])
|
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
"""@title: Cameras
|
|
2
|
+
|
|
3
|
+
Camera hardware specifications and image observations.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
import jaxtyping as jt
|
|
8
|
+
from enum import Enum
|
|
9
|
+
from dataclasses import dataclass
|
|
10
|
+
from beartype import beartype
|
|
11
|
+
|
|
12
|
+
from ..util.validation import check_hint
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class Camera_DeviceType(Enum):
|
|
16
|
+
Realsense_D405 = "realsense_d405"
|
|
17
|
+
"""Realsense D405 Stereo Camera. Both imagers are RGB, so one observation carries 2 views (left, right)."""
|
|
18
|
+
|
|
19
|
+
@property
|
|
20
|
+
def n_views(self) -> int:
|
|
21
|
+
"""Number of frames delivered per observation (1 for mono, 2 for a stereo pair)."""
|
|
22
|
+
return _CAMERA_N_VIEWS[self]
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
_CAMERA_N_VIEWS: dict[Camera_DeviceType, int] = {
|
|
26
|
+
Camera_DeviceType.Realsense_D405: 2,
|
|
27
|
+
}
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@dataclass(frozen=True)
|
|
31
|
+
class Camera_DeviceSpec:
|
|
32
|
+
device_type : Camera_DeviceType
|
|
33
|
+
"""Device type of the camera"""
|
|
34
|
+
|
|
35
|
+
resolution : tuple[int, int]
|
|
36
|
+
"""Resolution of each view of the camera stream, in tuple of (height, width), unit: [px]"""
|
|
37
|
+
|
|
38
|
+
refresh_rate : int
|
|
39
|
+
"""Refresh rate of the camera, unit: [fps]"""
|
|
40
|
+
|
|
41
|
+
jpeg_compression: int | None
|
|
42
|
+
"""jpeg compression quality value if specified, if None then no compression will be made."""
|
|
43
|
+
|
|
44
|
+
@property
|
|
45
|
+
def n_views(self) -> int:
|
|
46
|
+
return self.device_type.n_views
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
@jt.jaxtyped(typechecker=beartype)
|
|
50
|
+
@dataclass(frozen=True)
|
|
51
|
+
class Camera_Observation:
|
|
52
|
+
"""One decoded observation of a camera: ``b`` RGB views captured together (``b = 2`` for a stereo pair).
|
|
53
|
+
|
|
54
|
+
Frames are delivered already decoded regardless of ``jpeg_compression``. Construction enforces
|
|
55
|
+
``uint8`` BHWC with 3 channels; ``validate`` binds (b, h, w) to the camera spec.
|
|
56
|
+
"""
|
|
57
|
+
|
|
58
|
+
timestamp: float
|
|
59
|
+
"""Monotonic capture time on the station clock, unit: [s]. Shared by all views."""
|
|
60
|
+
|
|
61
|
+
rgb : jt.UInt8[np.ndarray, "b h w 3"]
|
|
62
|
+
"""Views in device order (stereo: index 0 = left, 1 = right)."""
|
|
63
|
+
|
|
64
|
+
@property
|
|
65
|
+
def n_views(self) -> int:
|
|
66
|
+
return int(self.rgb.shape[0])
|
|
67
|
+
|
|
68
|
+
def validate(self, spec: Camera_DeviceSpec) -> None:
|
|
69
|
+
h, w = spec.resolution
|
|
70
|
+
check_hint("rgb", self.rgb, jt.UInt8[np.ndarray, f"{spec.n_views} {h} {w} 3"])
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""@title: Evaluation Requests
|
|
2
|
+
|
|
3
|
+
What a customer submits to a station: the task, the policy, and how many
|
|
4
|
+
episodes to run. The station executes the request as uploaded; it does not
|
|
5
|
+
get to say what is being evaluated."""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from dataclasses import dataclass
|
|
10
|
+
|
|
11
|
+
from .task import Task_Spec
|
|
12
|
+
|
|
13
|
+
from .policy import Policy
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
@dataclass(frozen=True)
|
|
17
|
+
class JobRequest:
|
|
18
|
+
task : Task_Spec
|
|
19
|
+
policy : Policy
|
|
20
|
+
episodes: int
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
"""@title: Robot Policies
|
|
2
|
+
|
|
3
|
+
## Lifecycle
|
|
4
|
+
|
|
5
|
+
`se3labs.interface.eval.policy.Policy` defines the station-driven episode lifecycle:
|
|
6
|
+
|
|
7
|
+
| Method | Called when | Return value |
|
|
8
|
+
| --- | --- | --- |
|
|
9
|
+
| `initialize(station, task)` | Before each episode | `Station_JointPose`: the pose the arms move to before rollout |
|
|
10
|
+
| `infer(observation)` | On each observation tick | `Station_ActionChunk`, or `None` to continue the current chunk |
|
|
11
|
+
| `reset()` | After an episode | Nothing; clear per-episode state |
|
|
12
|
+
|
|
13
|
+
Model initialization and access to `task.instruction` belong in `initialize()`.
|
|
14
|
+
Inference frequency, preprocessing, and chunk length are controlled by the policy.
|
|
15
|
+
|
|
16
|
+
## Observations and actions
|
|
17
|
+
|
|
18
|
+
`station.arms` and `station.cameras` describe the available hardware. Observation
|
|
19
|
+
and action dictionaries use those same names. A bimanual station uses `left`
|
|
20
|
+
and `right` for its arms; a single-arm station uses `arm`.
|
|
21
|
+
|
|
22
|
+
Joint targets are absolute angles in radians. Gripper positions are open
|
|
23
|
+
fractions: `0` is closed, `1` is open. Each arm action contains arrays shaped
|
|
24
|
+
`(horizon, n_arm_dof)` and `(horizon, n_gripper_dof)`.
|
|
25
|
+
|
|
26
|
+
An action chunk's row `i` targets station time `timestamp + i / rate`. Use the
|
|
27
|
+
observation's station-clock timestamps, not your computer's wall clock.
|
|
28
|
+
All arms in a chunk share its horizon and rate. A new chunk replaces the
|
|
29
|
+
remaining targets from the previous chunk; returning `None` keeps that chunk.
|
|
30
|
+
|
|
31
|
+
## Task specification
|
|
32
|
+
|
|
33
|
+
A `Task_Spec` contains an operator-facing `description` and `reset_instruction`,
|
|
34
|
+
a model-facing `instruction`, a `max_rollout_time` in seconds, and a
|
|
35
|
+
`Task_ScoreSheet`. Its required binary `success` question determines whether
|
|
36
|
+
the episode passed. Additional questions can be binary or multiple choice.
|
|
37
|
+
The station operator supplies the answers.
|
|
38
|
+
|
|
39
|
+
## Episode completion and errors
|
|
40
|
+
|
|
41
|
+
Raising `se3labs.interface.eval.policy.EndEpisode` from `infer()` requests
|
|
42
|
+
normal scoring, reset, and homing. Other inference exceptions are added to
|
|
43
|
+
`SessionReport.errors`; that tick sends no new chunk. An initialization
|
|
44
|
+
exception ends the session and propagates to the caller.
|
|
45
|
+
|
|
46
|
+
For a complete implementation, follow the
|
|
47
|
+
[Python API tutorial](../../../tutorials/evaluate-your-policy/python-api.md). Data types are listed in the
|
|
48
|
+
[Python Interface reference](../../interface.md), and session reports in the
|
|
49
|
+
[Python SDK reference](../../sdk.md).
|
|
50
|
+
"""
|
|
51
|
+
from abc import ABC, abstractmethod
|
|
52
|
+
|
|
53
|
+
from .station import Station_Spec, Station_Observation, Station_ActionChunk, Station_JointPose
|
|
54
|
+
from .task import Task_Spec
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class EndEpisode(Exception):
|
|
58
|
+
"""Raise from ``infer`` to finish normally: scoring, reset, then station homing."""
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class Policy(ABC):
|
|
62
|
+
@abstractmethod
|
|
63
|
+
def initialize(self, station: Station_Spec, task: Task_Spec) -> Station_JointPose:
|
|
64
|
+
"""Prepare for one episode of ``task`` on ``station`` and return the pose the arms start from."""
|
|
65
|
+
|
|
66
|
+
@abstractmethod
|
|
67
|
+
def infer(self, observation: Station_Observation) -> Station_ActionChunk | None:
|
|
68
|
+
"""A new chunk, None to continue, or raise EndEpisode to score and home."""
|
|
69
|
+
|
|
70
|
+
def reset(self) -> None:
|
|
71
|
+
"""Clear per-episode state. Called once after the episode ends. No-op by default."""
|
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
"""@title: Station Repository
|
|
2
|
+
|
|
3
|
+
Registry of known station specs.
|
|
4
|
+
"""
|
|
5
|
+
from dataclasses import replace
|
|
6
|
+
from . import arm, camera
|
|
7
|
+
from .station import Bimanual_Station_Spec, Station_BaseSpec, Teleop_Bimanual_Station_Spec
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
Bimanual_YAM_Station = Bimanual_Station_Spec(
|
|
11
|
+
base=Station_BaseSpec.Static,
|
|
12
|
+
arm =arm.MotorizedArmSpec(
|
|
13
|
+
arm =arm.MotorizedArm_ArmType.YAM_6DoF,
|
|
14
|
+
gripper =arm.MotorizedArm_GripperType.I2RT_linear_gripper_1DoF,
|
|
15
|
+
refresh_rate=100,
|
|
16
|
+
control_rate=100,
|
|
17
|
+
),
|
|
18
|
+
wrist_camera=camera.Camera_DeviceSpec(
|
|
19
|
+
device_type =camera.Camera_DeviceType.Realsense_D405,
|
|
20
|
+
resolution =(360, 640), # (h, w); 2x downsample for 720p.
|
|
21
|
+
refresh_rate =30,
|
|
22
|
+
jpeg_compression=80
|
|
23
|
+
),
|
|
24
|
+
topdown_camera=camera.Camera_DeviceSpec(
|
|
25
|
+
device_type =camera.Camera_DeviceType.Realsense_D405,
|
|
26
|
+
resolution =(360, 640), # (h, w); 2x downsample for 720p.
|
|
27
|
+
refresh_rate =30,
|
|
28
|
+
jpeg_compression=80
|
|
29
|
+
)
|
|
30
|
+
)
|
|
31
|
+
"""Static bimanual YAM station: two YAM 6-DoF arms with I2RT grippers, two wrist D405s, one top-down D405."""
|
|
32
|
+
|
|
33
|
+
Teleop_Bimanual_YAM_Station = Teleop_Bimanual_Station_Spec(
|
|
34
|
+
base =Bimanual_YAM_Station.base,
|
|
35
|
+
arm =Bimanual_YAM_Station.arm,
|
|
36
|
+
leader_arm =replace(Bimanual_YAM_Station.arm, gripper=arm.MotorizedArm_GripperType.I2RT_teaching_handle),
|
|
37
|
+
wrist_camera =Bimanual_YAM_Station.wrist_camera,
|
|
38
|
+
topdown_camera=Bimanual_YAM_Station.topdown_camera,
|
|
39
|
+
)
|
|
40
|
+
"""``Bimanual_YAM_Station`` with two official YAM teaching-handle leaders."""
|
|
41
|
+
|
|
42
|
+
Teleop_Bimanual_OpenYAM_Station = replace(Teleop_Bimanual_YAM_Station, leader_arm=Bimanual_YAM_Station.arm)
|
|
43
|
+
"""The same layout with motorized grippers on the OpenYAM leaders."""
|
|
@@ -0,0 +1,173 @@
|
|
|
1
|
+
"""@title: Stations
|
|
2
|
+
|
|
3
|
+
Station specs (hardware capability) and the station-level observation / action containers.
|
|
4
|
+
|
|
5
|
+
A station is a named collection of arms and cameras. Observations and actions are keyed by
|
|
6
|
+
those names so that a bimanual and a single-arm station share one container type.
|
|
7
|
+
|
|
8
|
+
The policy side of the contract is ``policy.Policy``: it consumes ``Station_Observation`` ticks and
|
|
9
|
+
produces ``Station_ActionChunk``s. Anything about how the policy runs (its rate, chunk length,
|
|
10
|
+
replanning, image preprocessing) is decided inside the policy, not declared to the SDK.
|
|
11
|
+
"""
|
|
12
|
+
from abc import ABC, abstractmethod
|
|
13
|
+
from enum import Enum
|
|
14
|
+
from dataclasses import dataclass
|
|
15
|
+
|
|
16
|
+
from . import arm as arm_types
|
|
17
|
+
from . import camera
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class Station_BaseSpec(Enum):
|
|
21
|
+
Static = "static"
|
|
22
|
+
"""Static manipulation base (the arm is mounted on a fixed, non-mobile base)."""
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class Station_Spec(ABC):
|
|
26
|
+
base: Station_BaseSpec
|
|
27
|
+
|
|
28
|
+
@property
|
|
29
|
+
@abstractmethod
|
|
30
|
+
def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
|
|
31
|
+
"""Arm name -> spec. Names are the keys of ``Station_Observation.arms`` and ``Station_ActionChunk.arms``."""
|
|
32
|
+
|
|
33
|
+
@property
|
|
34
|
+
@abstractmethod
|
|
35
|
+
def cameras(self) -> dict[str, camera.Camera_DeviceSpec]:
|
|
36
|
+
"""Camera name -> spec. Names are the keys of ``Station_Observation.cameras``."""
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass(frozen=True)
|
|
40
|
+
class Bimanual_Station_Spec(Station_Spec):
|
|
41
|
+
base : Station_BaseSpec
|
|
42
|
+
|
|
43
|
+
arm : arm_types.MotorizedArmSpec
|
|
44
|
+
"""Spec shared by the ``left`` and ``right`` arms."""
|
|
45
|
+
|
|
46
|
+
wrist_camera : camera.Camera_DeviceSpec
|
|
47
|
+
"""Spec shared by the ``wrist_left`` and ``wrist_right`` cameras."""
|
|
48
|
+
|
|
49
|
+
topdown_camera: camera.Camera_DeviceSpec
|
|
50
|
+
|
|
51
|
+
@property
|
|
52
|
+
def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
|
|
53
|
+
return {"left": self.arm, "right": self.arm}
|
|
54
|
+
|
|
55
|
+
@property
|
|
56
|
+
def cameras(self) -> dict[str, camera.Camera_DeviceSpec]:
|
|
57
|
+
return {
|
|
58
|
+
"wrist_left" : self.wrist_camera,
|
|
59
|
+
"wrist_right": self.wrist_camera,
|
|
60
|
+
"topdown" : self.topdown_camera,
|
|
61
|
+
}
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
@dataclass(frozen=True)
|
|
65
|
+
class Teleop_Bimanual_Station_Spec(Bimanual_Station_Spec):
|
|
66
|
+
"""A bimanual station with a second pair of arms, ``leader_left`` and
|
|
67
|
+
``leader_right``, beside ``left`` and ``right``."""
|
|
68
|
+
|
|
69
|
+
leader_arm: arm_types.MotorizedArmSpec
|
|
70
|
+
|
|
71
|
+
@property
|
|
72
|
+
def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
|
|
73
|
+
return {"left": self.arm, "right": self.arm, "leader_left": self.leader_arm, "leader_right": self.leader_arm}
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
@dataclass(frozen=True)
|
|
77
|
+
class Single_Arm_Station_Spec(Station_Spec):
|
|
78
|
+
base : Station_BaseSpec
|
|
79
|
+
arm : arm_types.MotorizedArmSpec
|
|
80
|
+
wrist_camera : camera.Camera_DeviceSpec
|
|
81
|
+
|
|
82
|
+
@property
|
|
83
|
+
def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
|
|
84
|
+
return {"arm": self.arm}
|
|
85
|
+
|
|
86
|
+
@property
|
|
87
|
+
def cameras(self) -> dict[str, camera.Camera_DeviceSpec]:
|
|
88
|
+
return {"wrist": self.wrist_camera}
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
@dataclass(frozen=True)
|
|
92
|
+
class Station_Observation:
|
|
93
|
+
"""One policy-side observation: the latest tick of every arm and camera.
|
|
94
|
+
|
|
95
|
+
Arms and cameras run at different rates, so each entry carries its own timestamp; consumers that
|
|
96
|
+
need temporal alignment (world models, chunked policies) should use those rather than assume sync.
|
|
97
|
+
"""
|
|
98
|
+
|
|
99
|
+
arms : dict[str, arm_types.MotorizedArm_Observation]
|
|
100
|
+
cameras: dict[str, camera.Camera_Observation]
|
|
101
|
+
|
|
102
|
+
def validate(self, station: Station_Spec) -> None:
|
|
103
|
+
"""Check every arm / camera of ``station`` is present with the shapes its spec promises."""
|
|
104
|
+
for name, spec in station.arms.items():
|
|
105
|
+
self.arms[name].validate(spec)
|
|
106
|
+
for name, spec in station.cameras.items():
|
|
107
|
+
self.cameras[name].validate(spec)
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
@dataclass(frozen=True)
|
|
111
|
+
class Station_JointPose:
|
|
112
|
+
"""One absolute joint configuration for every arm of a station."""
|
|
113
|
+
|
|
114
|
+
arms: dict[str, arm_types.MotorizedArm_JointPose]
|
|
115
|
+
|
|
116
|
+
def validate(self, station: Station_Spec) -> None:
|
|
117
|
+
"""Check the pose names exactly the station's arms with the right shapes."""
|
|
118
|
+
if set(self.arms) != set(station.arms):
|
|
119
|
+
raise ValueError(f"pose arms {sorted(self.arms)} != station arms {sorted(station.arms)}")
|
|
120
|
+
for name, spec in station.arms.items():
|
|
121
|
+
try:
|
|
122
|
+
self.arms[name].validate(spec)
|
|
123
|
+
except (TypeError, ValueError) as e:
|
|
124
|
+
raise type(e)(f"arm '{name}': {e}") from None
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
@dataclass(frozen=True)
|
|
128
|
+
class Station_ActionChunk:
|
|
129
|
+
"""A chunk of ``horizon`` absolute joint-position targets for every arm of a station.
|
|
130
|
+
|
|
131
|
+
Row ``i`` of every arm is the target at ``timestamp + i / rate``. The station interpolates rows
|
|
132
|
+
up to each arm's ``control_rate`` and executes them until the policy returns a new chunk, which
|
|
133
|
+
replaces the remainder of the current one.
|
|
134
|
+
"""
|
|
135
|
+
|
|
136
|
+
timestamp: float
|
|
137
|
+
"""Station-clock time of the observation this chunk was predicted from, unit: [s]."""
|
|
138
|
+
|
|
139
|
+
rate : float
|
|
140
|
+
"""Spacing between consecutive rows, unit: [Hz]. Chosen by the policy per chunk."""
|
|
141
|
+
|
|
142
|
+
arms : dict[str, arm_types.MotorizedArm_ActionChunk]
|
|
143
|
+
|
|
144
|
+
def __post_init__(self) -> None:
|
|
145
|
+
if self.rate <= 0:
|
|
146
|
+
raise ValueError(f"rate must be positive, got {self.rate}")
|
|
147
|
+
horizons = {a.horizon for a in self.arms.values()}
|
|
148
|
+
if len(horizons) != 1:
|
|
149
|
+
raise ValueError(f"inconsistent horizons across arms: {horizons}")
|
|
150
|
+
|
|
151
|
+
@property
|
|
152
|
+
def horizon(self) -> int:
|
|
153
|
+
return next(iter(self.arms.values())).horizon
|
|
154
|
+
|
|
155
|
+
def __len__(self) -> int:
|
|
156
|
+
return self.horizon
|
|
157
|
+
|
|
158
|
+
def __getitem__(self, index: int) -> Station_JointPose:
|
|
159
|
+
"""Row ``index`` of the chunk, for every arm, as one station pose. Negative indices count from the end."""
|
|
160
|
+
return Station_JointPose({name: chunk[index] for name, chunk in self.arms.items()})
|
|
161
|
+
|
|
162
|
+
def validate(self, station: Station_Spec) -> None:
|
|
163
|
+
"""Check the chunk names exactly the station's arms, with the right shapes and a feasible rate."""
|
|
164
|
+
if set(self.arms) != set(station.arms):
|
|
165
|
+
raise ValueError(f"action arms {sorted(self.arms)} != station arms {sorted(station.arms)}")
|
|
166
|
+
h = self.horizon
|
|
167
|
+
for name, spec in station.arms.items():
|
|
168
|
+
if self.rate > spec.control_rate:
|
|
169
|
+
raise ValueError(f"arm '{name}': chunk rate {self.rate} Hz exceeds control_rate {spec.control_rate} Hz")
|
|
170
|
+
try:
|
|
171
|
+
self.arms[name].validate(spec, h)
|
|
172
|
+
except (TypeError, ValueError) as e:
|
|
173
|
+
raise type(e)(f"arm '{name}': {e}") from None
|