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.
Files changed (35) hide show
  1. se3labs_interface-0.0.1/PKG-INFO +14 -0
  2. se3labs_interface-0.0.1/pyproject.toml +37 -0
  3. se3labs_interface-0.0.1/setup.cfg +4 -0
  4. se3labs_interface-0.0.1/src/se3labs/interface/__init__.py +8 -0
  5. se3labs_interface-0.0.1/src/se3labs/interface/eval/__init__.py +40 -0
  6. se3labs_interface-0.0.1/src/se3labs/interface/eval/arm.py +194 -0
  7. se3labs_interface-0.0.1/src/se3labs/interface/eval/camera.py +70 -0
  8. se3labs_interface-0.0.1/src/se3labs/interface/eval/job.py +20 -0
  9. se3labs_interface-0.0.1/src/se3labs/interface/eval/policy.py +71 -0
  10. se3labs_interface-0.0.1/src/se3labs/interface/eval/repository.py +43 -0
  11. se3labs_interface-0.0.1/src/se3labs/interface/eval/station.py +173 -0
  12. se3labs_interface-0.0.1/src/se3labs/interface/eval/task.py +164 -0
  13. se3labs_interface-0.0.1/src/se3labs/interface/util/__init__.py +0 -0
  14. se3labs_interface-0.0.1/src/se3labs/interface/util/validation.py +16 -0
  15. se3labs_interface-0.0.1/src/se3labs/interface/wire/__init__.py +32 -0
  16. se3labs_interface-0.0.1/src/se3labs/interface/wire/codec.py +288 -0
  17. se3labs_interface-0.0.1/src/se3labs/interface/wire/job_status.proto +43 -0
  18. se3labs_interface-0.0.1/src/se3labs/interface/wire/job_status_pb2.py +41 -0
  19. se3labs_interface-0.0.1/src/se3labs/interface/wire/job_status_pb2.pyi +30 -0
  20. se3labs_interface-0.0.1/src/se3labs/interface/wire/job_status_pb2_grpc.py +102 -0
  21. se3labs_interface-0.0.1/src/se3labs/interface/wire/policy_session.proto +320 -0
  22. se3labs_interface-0.0.1/src/se3labs/interface/wire/policy_session_pb2.py +124 -0
  23. se3labs_interface-0.0.1/src/se3labs/interface/wire/policy_session_pb2.pyi +393 -0
  24. se3labs_interface-0.0.1/src/se3labs/interface/wire/policy_session_pb2_grpc.py +117 -0
  25. se3labs_interface-0.0.1/src/se3labs/interface/wire/station_link.proto +49 -0
  26. se3labs_interface-0.0.1/src/se3labs/interface/wire/station_link_pb2.py +48 -0
  27. se3labs_interface-0.0.1/src/se3labs/interface/wire/station_link_pb2.pyi +49 -0
  28. se3labs_interface-0.0.1/src/se3labs/interface/wire/station_link_pb2_grpc.py +107 -0
  29. se3labs_interface-0.0.1/src/se3labs_interface.egg-info/PKG-INFO +14 -0
  30. se3labs_interface-0.0.1/src/se3labs_interface.egg-info/SOURCES.txt +33 -0
  31. se3labs_interface-0.0.1/src/se3labs_interface.egg-info/dependency_links.txt +1 -0
  32. se3labs_interface-0.0.1/src/se3labs_interface.egg-info/requires.txt +10 -0
  33. se3labs_interface-0.0.1/src/se3labs_interface.egg-info/top_level.txt +1 -0
  34. se3labs_interface-0.0.1/tests/test_spec.py +224 -0
  35. 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,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
@@ -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